feat(omnivoice): voice prompts that survive restarts + opt-in FlashInfer (~2.2x) (#1565)
* feat(omnivoice): port upstream VoiceClonePrompt persistence + FlashInfer opt-in Upstream k2-fsa teardown ports, verified with generated voice samples: - VoiceClonePrompt.save()/.load() (upstream format v1, weights_only-safe) on the vendored model, and a disk layer under the in-memory prompt LRU (DATA_DIR/prompt_cache, keyed by ref path+mtime+ref_text+preprocess, 32 newest kept, OMNIVOICE_PROMPT_DISK_CACHE=0 opts out). First generation of a session with a known voice skips the reference re-encode and any auto-transcription pass — verified across two real processes (encodes=1 then encodes=0, same voice). - omnivoice_flashinfer.py ported (packed CFG attention, fused kernels, optional CUDA graphs), schedule adapted to our num_step+1 divergence. Opt-in via OMNIVOICE_FLASHINFER=1|graph, CUDA-only, replaces torch.compile for the session; missing package / apply failure / runtime failure all degrade with a named reason (same #278 contract as compile: classify → unapply → retry once, session latch). Measured 2.20x at batch=1 on an RTX 4090 with byte-identical text and clean ASR round-trip. - Docs: OmniVoice guide gains instruct+reference combination semantics (consistent instruct stabilizes cloning, reference wins conflicts), inline pronunciation control (pinyin / CMU), prompt persistence, and corrects the 'no voice design' claim; performance.md documents both new env knobs. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> * chore: point changelog entries at the real PR number (#1565) Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> * fix(pr): harden FlashInfer lifecycle + prompt-cache writes per review Bot harvest round 1 (#1565): unapply on apply-failure (half-patched model could crash the next render); pin eager-mode FlashInfer inference to one thread too — the attention plan and packed position ids are per-generation module state, so interleaved _gpu_pool workers would corrupt each other; restore the CAPTURED pre-apply attention impl (could be flash_attention_2) instead of assuming sdpa; unique tmp name per prompt-cache write; correct the _forward_logits layout docstring; resolve VoiceClonePrompt at test runtime; docs — Known limits keeps only the limitation, performance.md states the VRAM cost and scopes the fallback claim to classified kernel failures. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> * fix(pr): round-2 review — publish only a fully restored model, redact latch reason, tighten CPU-persistence test Greptile: the runtime fallback now unapplies BEFORE swapping generate, so a concurrent render keeps queuing behind the thread-affinity wrapper while teardown mutates modules. CodeRabbit: FlashInfer failure reasons pass through core.failure.sanitize before latching/logging (wheel paths embed the user's home); the save-portability test now creates the tokens on CUDA when available and asserts the persisted payload itself is CPU-resident. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> * fix(pr): fail-closed latch reason when the sanitizer itself breaks CodeQL empty-except + CodeRabbit round 3: if core.failure.sanitize raises, the raw reason (home paths, wheel paths) was latched anyway. Now only the exception class survives with a fixed redaction note; two regression tests (normal redaction + sanitizer failure). Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> --------- Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
4db02d0c97
commit
2d5f2e800e
@@ -18,12 +18,15 @@ the frozen-backend fallback mirror it for their toolchains.
|
||||
- The backend binds its port immediately and reports startup progress live — `/health` answers 503-with-step and a new `/startup/progress` endpoint lists every step while PyTorch, API routes, and database migrations load in the background, so "starting at step X" is never mistakable for "dead"; the desktop splash narrates each step (#1550)
|
||||
|
||||
### Added
|
||||
- Voices you've cloned stay "warm" across restarts — encoded references now persist to disk (~10 KB each), so the first generation of a session skips the re-encode and any transcription pass; `OMNIVOICE_PROMPT_DISK_CACHE=0` opts out (#1565)
|
||||
- Optional FlashInfer acceleration for the default engine on CUDA (`OMNIVOICE_FLASHINFER=1`, ~2.2x measured) — needs the optional `flashinfer-python` package; missing package or kernel failure logs why and falls back to the standard path (#1565)
|
||||
- The bug reporter notices when you're on an outdated build and offers the latest release before filing — with a "File anyway" escape hatch — and stamps a `Build status` line into every report so up-to-date reports are tellable from stale ones (#1547)
|
||||
- Settings → Performance & Device gains a compute-device override (Auto / CUDA / ROCm / XPU / MPS / CPU, or `OMNIVOICE_DEVICE`) — pin the device when auto-detect picks wrong; only devices your machine actually has are offered (#1557)
|
||||
|
||||
### Docs
|
||||
- The READMEs now lead with download buttons and a three-step first-clone walkthrough, and a new benchmarks page anchors measured per-engine/per-device numbers on the in-repo harness (#1555)
|
||||
- Every engine now has its own guide — 21 new pages under docs/engines plus an index covering all 16 TTS and 11 ASR engines, linked from both READMEs (#1556)
|
||||
- The OmniVoice guide now covers combining style attributes with a reference clip (consistent instruct stabilizes cloning; the reference wins conflicts), inline pronunciation control (pinyin / CMU phonemes), and corrects the claim that the default engine can't do voice design — it can, from attributes (#1565)
|
||||
|
||||
### Fixed
|
||||
- The crash-isolated ASR sidecar and its download preflight now agree on which model to load — setting the shared faster-whisper model variable applies to both variants instead of the sidecar quietly using a different one (#1556)
|
||||
|
||||
@@ -57,6 +57,99 @@ def _force_compile_requested() -> bool:
|
||||
return value.strip().lower() in {"1", "true", "yes", "on"}
|
||||
|
||||
|
||||
# ── FlashInfer opt-in (upstream k2-fsa port) ────────────────────────────────
|
||||
# Explicit power-user opt-in, CUDA-only: OMNIVOICE_FLASHINFER=1 patches the
|
||||
# OmniVoice model with flashinfer packed attention (~2x per upstream's
|
||||
# benchmarks); =graph additionally captures CUDA graphs (best at batch=1).
|
||||
# Off by default — `flashinfer` is not a shipped dependency, and an
|
||||
# optimization must never be a point of failure. Session-sticky failure
|
||||
# latch mirrors torch.compile's (#278).
|
||||
_FLASHINFER_ENV = "OMNIVOICE_FLASHINFER"
|
||||
_flashinfer_runtime_failure: Optional[str] = None
|
||||
|
||||
|
||||
def flashinfer_mode() -> str:
|
||||
"""The user's ``OMNIVOICE_FLASHINFER`` request: 'off' | 'on' | 'graph'.
|
||||
|
||||
Unknown values normalize to 'off' with a log line naming the env var, so
|
||||
a typo degrades to the default path instead of half-applying.
|
||||
"""
|
||||
value = os.environ.get(_FLASHINFER_ENV, "").strip().lower()
|
||||
if value in {"", "0", "false", "no", "off"}:
|
||||
return "off"
|
||||
if value in {"1", "true", "yes", "on"}:
|
||||
return "on"
|
||||
if value == "graph":
|
||||
return "graph"
|
||||
logger.warning(
|
||||
"%s=%r not recognized (valid: 0, 1, graph) — FlashInfer stays off.",
|
||||
_FLASHINFER_ENV, value,
|
||||
)
|
||||
return "off"
|
||||
|
||||
|
||||
def should_flashinfer(device: str) -> str:
|
||||
"""Resolve the FlashInfer request against this host: 'off' | 'on' | 'graph'.
|
||||
|
||||
Requires all of: the ``OMNIVOICE_FLASHINFER`` opt-in, device == "cuda"
|
||||
(flashinfer is CUDA-only), the ``flashinfer`` package importable, and no
|
||||
earlier runtime failure this session. Every refusal is logged with the
|
||||
reason and the knob's name — the user asked for it, so silence would read
|
||||
as "the setting doesn't work".
|
||||
"""
|
||||
mode = flashinfer_mode()
|
||||
if mode == "off":
|
||||
return "off"
|
||||
if device != "cuda":
|
||||
logger.warning(
|
||||
"%s requested but the compute device is %r — FlashInfer is "
|
||||
"CUDA-only, continuing without it.", _FLASHINFER_ENV, device,
|
||||
)
|
||||
return "off"
|
||||
if importlib.util.find_spec("flashinfer") is None:
|
||||
logger.warning(
|
||||
"%s requested but the `flashinfer` package is not installed — "
|
||||
"continuing without it. Install with: uv pip install "
|
||||
"flashinfer-python flashinfer-jit-cache "
|
||||
"--extra-index-url https://flashinfer.ai/whl/cu128/ "
|
||||
"(pick the index matching your CUDA build).", _FLASHINFER_ENV,
|
||||
)
|
||||
return "off"
|
||||
if _flashinfer_runtime_failure is not None:
|
||||
logger.info(
|
||||
"FlashInfer skipped: failed earlier this session (%s) — using the "
|
||||
"standard path.", _flashinfer_runtime_failure,
|
||||
)
|
||||
return "off"
|
||||
return mode
|
||||
|
||||
|
||||
def mark_flashinfer_runtime_failure(reason: str) -> None:
|
||||
"""Latch a FlashInfer apply/runtime failure for the rest of the process,
|
||||
same contract as ``mark_compile_runtime_failure``."""
|
||||
global _flashinfer_runtime_failure
|
||||
try:
|
||||
# Import/kernel errors embed absolute paths (wheels under the user's
|
||||
# home) — redact before latching, since the reason is logged here and
|
||||
# re-logged on every later skip.
|
||||
from core.failure import sanitize
|
||||
|
||||
reason = sanitize(reason)
|
||||
except Exception:
|
||||
# Fail closed: if the redactor itself breaks, latching the raw text
|
||||
# would defeat the redaction. Keep only the exception class (the part
|
||||
# before ':' in our "Type: message" reasons) and drop the message.
|
||||
reason = (
|
||||
f"{(reason or '').split(':', 1)[0][:80]} "
|
||||
"(details redacted: sanitizer unavailable)"
|
||||
).strip()
|
||||
_flashinfer_runtime_failure = reason or "unknown FlashInfer runtime failure"
|
||||
logger.warning(
|
||||
"FlashInfer disabled for this session after a runtime failure: %s",
|
||||
_flashinfer_runtime_failure,
|
||||
)
|
||||
|
||||
|
||||
def _cuda_arch_supported_for_compile() -> "tuple[bool, str]":
|
||||
"""Check the GPU's architecture against this torch build's arch list.
|
||||
|
||||
|
||||
@@ -1376,6 +1376,122 @@ def _install_compile_fallback(_model) -> None:
|
||||
_model.generate = _generate_with_compile_fallback
|
||||
|
||||
|
||||
# ── FlashInfer runtime fallback (upstream k2-fsa port) ──────────────────────
|
||||
|
||||
|
||||
def _is_flashinfer_runtime_failure(exc: BaseException) -> bool:
|
||||
"""True when an exception originates in the FlashInfer fast path (the
|
||||
flashinfer package, our omnivoice_flashinfer patch module, or CUDA-graph
|
||||
capture/replay) rather than in the model or the request itself. Same
|
||||
chain/traceback walk as ``_is_compile_runtime_failure``."""
|
||||
import traceback as _tb
|
||||
|
||||
tb_markers = ("/flashinfer/", "omnivoice_flashinfer")
|
||||
msg_markers = ("flashinfer", "cuda graph", "cudagraph")
|
||||
seen: set[int] = set()
|
||||
cur: BaseException | None = exc
|
||||
while cur is not None and id(cur) not in seen:
|
||||
seen.add(id(cur))
|
||||
mod = type(cur).__module__ or ""
|
||||
if mod.startswith("flashinfer"):
|
||||
return True
|
||||
msg = str(cur).lower()
|
||||
if any(marker in msg for marker in msg_markers):
|
||||
return True
|
||||
try:
|
||||
for frame in _tb.extract_tb(cur.__traceback__):
|
||||
filename = (frame.filename or "").replace("\\", "/")
|
||||
if any(marker in filename for marker in tb_markers):
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
if cur.__cause__ is not None:
|
||||
cur = cur.__cause__
|
||||
elif not cur.__suppress_context__:
|
||||
cur = cur.__context__
|
||||
else:
|
||||
cur = None
|
||||
return False
|
||||
|
||||
|
||||
def _unapply_flashinfer(_model) -> None:
|
||||
"""Restore the standard execution path on a FlashInfer-patched model.
|
||||
|
||||
``apply_flashinfer`` works entirely through *instance-level* state —
|
||||
MethodType-bound ``forward``/``_generate_iterative`` overrides and
|
||||
``_fi_*`` attributes — so deleting those attributes restores the class
|
||||
implementations exactly. The attention implementation is restored to the
|
||||
one captured before apply (``_fi_orig_attn_impl`` — could be
|
||||
flash_attention_2, not just sdpa), and use_cache is re-enabled."""
|
||||
llm = getattr(_model, "llm", None)
|
||||
orig_attn = getattr(_model, "_fi_orig_attn_impl", None) or "sdpa"
|
||||
if llm is not None:
|
||||
for module in llm.modules():
|
||||
if "forward" in vars(module):
|
||||
del module.forward
|
||||
for attr in ("_fi_w_qkv", "_fi_qkv_split", "_fi_rope_theta", "_fi_w_gate_up"):
|
||||
if attr in vars(module):
|
||||
delattr(module, attr)
|
||||
try:
|
||||
llm.set_attn_implementation(orig_attn)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"failed to restore %s attention after FlashInfer", orig_attn
|
||||
)
|
||||
llm.config.use_cache = True
|
||||
for attr in (
|
||||
"_fi_orig_attn_impl",
|
||||
"_generate_iterative",
|
||||
"_fi_runner",
|
||||
"_fi_graph_cache",
|
||||
"_fi_enable_cuda_graph",
|
||||
"_fi_graph_buckets",
|
||||
"_fi_overhead_budget",
|
||||
):
|
||||
if attr in vars(_model):
|
||||
delattr(_model, attr)
|
||||
|
||||
|
||||
def _install_flashinfer_fallback(_model) -> None:
|
||||
"""Wrap ``model.generate`` so a FlashInfer failure at inference time falls
|
||||
back to the standard path instead of failing the generation — the same
|
||||
contract as ``_install_compile_fallback`` (#278): an optimization must
|
||||
never turn a working generation into an error."""
|
||||
orig_generate = _model.generate
|
||||
|
||||
def _generate_with_flashinfer_fallback(*args, **kwargs):
|
||||
try:
|
||||
return orig_generate(*args, **kwargs)
|
||||
except Exception as exc:
|
||||
if not _is_flashinfer_runtime_failure(exc):
|
||||
raise
|
||||
logger.warning(
|
||||
"FlashInfer runtime failure during generation (%s: %s) — "
|
||||
"restoring the standard path and disabling FlashInfer for "
|
||||
"this session. Generation is being retried without it.",
|
||||
type(exc).__name__, exc,
|
||||
)
|
||||
from services import engine_env
|
||||
engine_env.mark_flashinfer_runtime_failure(
|
||||
f"{type(exc).__name__}: {exc}"
|
||||
)
|
||||
# Unapply BEFORE exposing the eager path: while the teardown
|
||||
# mutates modules, _model.generate still routes through the
|
||||
# thread-affinity wrapper, so a concurrent render queues behind
|
||||
# this call instead of racing the half-restored model (Greptile,
|
||||
# #1565 round 2). Only a fully restored model is published.
|
||||
_unapply_flashinfer(_model)
|
||||
_model.generate = orig_generate
|
||||
try:
|
||||
return orig_generate(*args, **kwargs)
|
||||
except Exception as plain_exc:
|
||||
# `from None`: a genuine standard-path failure must not be
|
||||
# chained to — and misread as — the FlashInfer error.
|
||||
raise plain_exc from None
|
||||
|
||||
_model.generate = _generate_with_flashinfer_fallback
|
||||
|
||||
|
||||
# ── #315: thread affinity for cudagraph-compiled models ─────────────────────
|
||||
# `torch.compile(mode="reduce-overhead")` captures CUDA graphs, and captured
|
||||
# graph state is **thread-local** (torch/_inductor/cudagraph_trees keys its
|
||||
@@ -2117,6 +2233,57 @@ def _load_model_sync():
|
||||
"to stop preloading it alongside TTS."
|
||||
) from asr_exc
|
||||
|
||||
# FlashInfer opt-in (upstream k2-fsa port): packed CFG attention +
|
||||
# fused kernels, ~2x on upstream's benchmarks. Applied INSTEAD of
|
||||
# torch.compile — both rewrite the llm's execution and they do not
|
||||
# compose. Best-effort: any apply failure latches the session off and
|
||||
# the standard path continues untouched.
|
||||
flashinfer_applied = False
|
||||
try:
|
||||
from services.engine_env import (
|
||||
mark_flashinfer_runtime_failure,
|
||||
should_flashinfer,
|
||||
)
|
||||
|
||||
fi_mode = should_flashinfer(device)
|
||||
if fi_mode != "off":
|
||||
_set_loading("compiling", "Applying FlashInfer kernels…")
|
||||
try:
|
||||
from omnivoice.models.omnivoice_flashinfer import apply_flashinfer
|
||||
|
||||
# Captured BEFORE apply so unapply (either the failure
|
||||
# branch below or the generate-time fallback) restores
|
||||
# the true prior implementation.
|
||||
_model._fi_orig_attn_impl = getattr(
|
||||
_model.llm.config, "_attn_implementation", "sdpa"
|
||||
)
|
||||
apply_flashinfer(_model, enable_cuda_graph=(fi_mode == "graph"))
|
||||
except Exception as fi_exc: # noqa: BLE001 — perf opt, never fatal
|
||||
mark_flashinfer_runtime_failure(
|
||||
f"{type(fi_exc).__name__}: {fi_exc}"
|
||||
)
|
||||
# apply_flashinfer mutates the model as it goes — a
|
||||
# failure partway leaves half-patched modules that would
|
||||
# crash the next render (Greptile, #1565). Restore fully.
|
||||
_unapply_flashinfer(_model)
|
||||
else:
|
||||
flashinfer_applied = True
|
||||
_install_flashinfer_fallback(_model)
|
||||
# BOTH modes pin inference to one thread. Graph mode for
|
||||
# the #315 reason (captured CUDA-graph state is
|
||||
# thread-local); eager mode because the FlashInfer
|
||||
# attention wrapper and packed position ids are planned
|
||||
# per generation in module state — two _gpu_pool workers
|
||||
# interleaving plan() and run() would corrupt each
|
||||
# other's layout (CodeRabbit/Greptile, #1565).
|
||||
_install_compile_thread_affinity(_model)
|
||||
logger.info(
|
||||
"FlashInfer applied (mode=%s) — torch.compile skipped "
|
||||
"for this load.", fi_mode,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("FlashInfer opt-in check failed; continuing without")
|
||||
|
||||
try:
|
||||
# plan-02 (#65): gate on Triton availability (+ user setting), not
|
||||
# just device==cuda. Triton has no Windows wheel, so the old
|
||||
@@ -2124,7 +2291,7 @@ def _load_model_sync():
|
||||
# falls back to eager there.
|
||||
from services.engine_env import should_torch_compile
|
||||
|
||||
if should_torch_compile(device):
|
||||
if not flashinfer_applied and should_torch_compile(device):
|
||||
_set_loading("compiling", "Compiling model (torch.compile)…")
|
||||
try:
|
||||
_model.llm = torch.compile(_model.llm, mode=_TORCH_COMPILE_MODE)
|
||||
|
||||
@@ -412,6 +412,95 @@ _PROMPT_CACHE_MAX = 8
|
||||
_prompt_cache: "OrderedDict[tuple, object]" = OrderedDict()
|
||||
_prompt_cache_lock = threading.Lock()
|
||||
|
||||
# Disk layer under the in-memory LRU (upstream k2-fsa VoiceClonePrompt.save/
|
||||
# load format). The in-memory cache dies with the process, so the first
|
||||
# generation of every session re-encodes each voice (~0.4 s + an ASR pass when
|
||||
# ref_text is missing). Encoded prompts are tiny (a (8, T) int token tensor +
|
||||
# transcript), so we persist them and reload across restarts. Keyed by the
|
||||
# same tuple as the memory cache — the ref file's mtime is inside the key, so
|
||||
# an edited reference never matches a stale file; stale files age out via the
|
||||
# mtime prune. Best-effort like the memory cache: any failure means "no disk
|
||||
# hit / no disk write", never a failed generation. OMNIVOICE_PROMPT_DISK_CACHE=0
|
||||
# disables the layer entirely.
|
||||
_PROMPT_DISK_CACHE_MAX = 32
|
||||
|
||||
|
||||
def _prompt_disk_dir():
|
||||
"""Return the prompt-cache directory (created on first use), or None when
|
||||
the layer is disabled or the directory can't be created."""
|
||||
if os.environ.get("OMNIVOICE_PROMPT_DISK_CACHE", "1") == "0":
|
||||
return None
|
||||
try:
|
||||
from core.config import DATA_DIR
|
||||
|
||||
path = os.path.join(str(DATA_DIR), "prompt_cache")
|
||||
os.makedirs(path, exist_ok=True)
|
||||
return path
|
||||
except Exception as e: # noqa: BLE001 — cache layer must never break synthesis
|
||||
logger.debug("prompt disk cache unavailable: %s", e)
|
||||
return None
|
||||
|
||||
|
||||
def _prompt_disk_path(cache_dir: str, key: tuple) -> str:
|
||||
import hashlib
|
||||
|
||||
digest = hashlib.sha256(repr(key).encode("utf-8")).hexdigest()[:32]
|
||||
return os.path.join(cache_dir, f"{digest}.pt")
|
||||
|
||||
|
||||
def _prompt_disk_load(key: tuple):
|
||||
"""Load a persisted prompt for ``key``, or None. Never raises."""
|
||||
cache_dir = _prompt_disk_dir()
|
||||
if cache_dir is None:
|
||||
return None
|
||||
path = _prompt_disk_path(cache_dir, key)
|
||||
if not os.path.exists(path):
|
||||
return None
|
||||
try:
|
||||
from omnivoice.models.omnivoice import VoiceClonePrompt
|
||||
|
||||
prompt = VoiceClonePrompt.load(path)
|
||||
# Freshen so the LRU prune (by mtime) keeps actively used voices.
|
||||
os.utime(path, None)
|
||||
return prompt
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("failed to load cached voice prompt %s: %s", path, e)
|
||||
try:
|
||||
os.remove(path) # corrupt/incompatible file — don't retry it forever
|
||||
except OSError:
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
def _prompt_disk_save(key: tuple, prompt) -> None:
|
||||
"""Persist ``prompt`` under ``key`` and prune old entries. Never raises."""
|
||||
cache_dir = _prompt_disk_dir()
|
||||
if cache_dir is None:
|
||||
return
|
||||
path = _prompt_disk_path(cache_dir, key)
|
||||
try:
|
||||
# Unique per write: two GPU-pool threads missing the same key must not
|
||||
# interleave writes into one tmp file (os.replace stays atomic).
|
||||
import uuid
|
||||
|
||||
tmp = f"{path}.tmp.{os.getpid()}.{uuid.uuid4().hex[:8]}"
|
||||
prompt.save(tmp)
|
||||
os.replace(tmp, path)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("failed to persist voice prompt to %s: %s", path, e)
|
||||
return
|
||||
try:
|
||||
entries = [
|
||||
os.path.join(cache_dir, f)
|
||||
for f in os.listdir(cache_dir)
|
||||
if f.endswith(".pt")
|
||||
]
|
||||
entries.sort(key=lambda p: os.path.getmtime(p), reverse=True)
|
||||
for old in entries[_PROMPT_DISK_CACHE_MAX:]:
|
||||
os.remove(old)
|
||||
except OSError as e:
|
||||
logger.debug("prompt disk cache prune skipped: %s", e)
|
||||
|
||||
|
||||
def _clone_prompt_key(ref_audio: str, ref_text, preprocess_prompt: bool = True):
|
||||
try:
|
||||
@@ -450,15 +539,24 @@ def _get_clone_prompt(
|
||||
if hit is not None:
|
||||
_prompt_cache.move_to_end(key)
|
||||
return hit
|
||||
try:
|
||||
# Encode outside the lock (slow). Mirrors exactly what generate() would
|
||||
# do inline for this ref (omnivoice.py:964-978), so output is identical.
|
||||
prompt = model.create_voice_clone_prompt(
|
||||
ref_audio, ref_text=ref_text, preprocess_prompt=preprocess_prompt
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 — fall back, never break synthesis
|
||||
logger.warning("voice-clone prompt precompute failed; using inline ref: %s", e)
|
||||
return None
|
||||
# Memory miss → disk (survives restarts). A disk hit skips the encode AND
|
||||
# the ASR transcription pass a ref_text-less reference would trigger.
|
||||
prompt = _prompt_disk_load(key)
|
||||
if prompt is None:
|
||||
try:
|
||||
# Encode outside the lock (slow). Mirrors exactly what generate()
|
||||
# would do inline for this ref (omnivoice.py:964-978), so output is
|
||||
# identical.
|
||||
prompt = model.create_voice_clone_prompt(
|
||||
ref_audio, ref_text=ref_text, preprocess_prompt=preprocess_prompt
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 — fall back, never break synthesis
|
||||
logger.warning(
|
||||
"voice-clone prompt precompute failed; using inline ref: %s", e
|
||||
)
|
||||
return None
|
||||
if store:
|
||||
_prompt_disk_save(key, prompt)
|
||||
if not store:
|
||||
return prompt
|
||||
with _prompt_cache_lock:
|
||||
|
||||
@@ -47,13 +47,39 @@ The env var overrides the persisted UI choice.
|
||||
co-loaded for the cloning path.
|
||||
- Output is 24 kHz mono; the shared mastering chain (highpass + compressor)
|
||||
is tuned for this rate and applied automatically.
|
||||
- Cloning takes a short reference clip (`ref_audio`); an optional transcript
|
||||
of the clip improves conditioning.
|
||||
- Cloning takes a short reference clip (`ref_audio`); 3–10 seconds is the
|
||||
sweet spot. A transcript of the clip improves conditioning — if the profile
|
||||
has none, VoiceStudio transcribes the clip automatically on first use and
|
||||
saves the result to the profile.
|
||||
- Encoded voice references persist on disk (`prompt_cache/` in the app data
|
||||
dir), so the first generation with a known voice after a restart skips the
|
||||
re-encode and any transcription pass. Set `OMNIVOICE_PROMPT_DISK_CACHE=0`
|
||||
to keep the cache in memory only.
|
||||
- Style attributes (`instruct`) and a reference clip can be **combined**:
|
||||
when they agree, the instruct stabilizes cloning for the attributes it
|
||||
names (upstream documents dialect cloning as the canonical case — dialect
|
||||
reference + matching dialect instruct). When they conflict, the reference
|
||||
audio wins.
|
||||
- Inline pronunciation control: Chinese via pinyin with tone numbers
|
||||
(`打ZHE2出售`), English via bracketed CMU phonemes (`[B EY1 S]`). Non-verbal
|
||||
tags like `[laughter]` are covered in
|
||||
[expressive-speech.md](../expressive-speech.md).
|
||||
- Voice design works from attributes (gender, age, pitch, whisper, English
|
||||
accents, Chinese dialects) via the Design tab — no reference audio needed.
|
||||
- Optional FlashInfer acceleration on CUDA: set `OMNIVOICE_FLASHINFER=1`
|
||||
(or `=graph` for CUDA-graph capture, best for one render at a time) after
|
||||
installing the `flashinfer-python` package — see
|
||||
[performance.md](../performance.md). Off by default; if the package is
|
||||
missing or a kernel fails, the app logs why and continues on the standard
|
||||
path.
|
||||
|
||||
## Known limits
|
||||
|
||||
- No voice design from a text description — use [VoxCPM2](voxcpm2.md) for
|
||||
that.
|
||||
- Voice design understands only the fixed attribute vocabulary — free-form
|
||||
design *prose* is mapped onto those attributes, and wording outside them
|
||||
is ignored. Design is trained on English and Chinese and can be unstable
|
||||
in low-resource languages; for description-driven design in other cases
|
||||
try [VoxCPM2](voxcpm2.md).
|
||||
- Below the 6 GB VRAM floor, expect very slow renders or budget timeouts;
|
||||
prefer [OmniVoice GGUF](omnivoice-gguf.md) or a CPU engine such as
|
||||
[PocketTTS](pockettts.md).
|
||||
|
||||
@@ -81,6 +81,8 @@ None of them are required — the defaults are chosen for the common case.
|
||||
| Variable | Default | What it does |
|
||||
|---|---|---|
|
||||
| `OMNIVOICE_DEVICE` | `auto` | Pin the compute device (`cuda` / `rocm` / `xpu` / `mps` / `cpu`) instead of auto-detect. Same control lives in **Settings → Performance & Device** (the env var wins over the UI pick). Honored only for devices the host actually has — a family that isn't detected is noted and ignored, never obeyed blindly. Applies at the next backend start. |
|
||||
| `OMNIVOICE_FLASHINFER` | `0` | CUDA-only accelerated decoding for the default engine via [FlashInfer](https://github.com/flashinfer-ai/flashinfer) kernels (packed CFG attention, fused RMSNorm/RoPE/GEMM) — ~2x on upstream's benchmarks. `1` enables it; `graph` also captures CUDA graphs (best when you render one thing at a time). Requires installing the optional `flashinfer-python` package into the backend environment first (`uv pip install flashinfer-python flashinfer-jit-cache --extra-index-url https://flashinfer.ai/whl/cu128/`, matching your CUDA build). Replaces `torch.compile` for that session, pins inference to a single GPU thread (the FlashInfer attention plan is per-generation state), and keeps fused copies of the attention/MLP weights resident (~roughly half the LLM's weight size extra VRAM) — leave it off on tight-VRAM cards. If the package is missing or a FlashInfer/CUDA-graph kernel fails at runtime, the app logs the reason and falls back to the standard path; failures outside those kernels (e.g. a genuine out-of-memory) surface normally. |
|
||||
| `OMNIVOICE_PROMPT_DISK_CACHE` | `1` | Persist encoded voice-clone references (`prompt_cache/` in the app data dir, ~10 KB per voice, 32 newest kept) so the first generation with a known voice after a restart skips the reference re-encode and any auto-transcription. Set `0` to keep the cache in memory only. |
|
||||
| `OMNIVOICE_IDLE_TIMEOUT_S` | `900` | Seconds of idle before the TTS model unloads to free memory. Raise it (e.g. `3600`) if you generate in bursts and dislike the ~8 s reload; lower it on tight-memory machines. |
|
||||
| `OMNIVOICE_SIDECAR_IDLE_TIMEOUT_S` | `300` | Same idea for sidecar engines (IndexTTS 2.5 etc.). |
|
||||
| `OMNIVOICE_LLM_CONCURRENCY` | `6` | Parallel LLM translation calls during a dub. Raise for a fast API endpoint, lower if your provider rate-limits. |
|
||||
|
||||
@@ -91,12 +91,60 @@ class OmniVoiceModelAssetError(RuntimeError):
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
_VOICE_CLONE_PROMPT_FORMAT_VERSION = 1
|
||||
|
||||
|
||||
@dataclass
|
||||
class VoiceClonePrompt:
|
||||
ref_audio_tokens: torch.Tensor # (C, T)
|
||||
ref_text: str
|
||||
ref_rms: float
|
||||
|
||||
def save(self, path: str) -> None:
|
||||
"""Save this prompt to ``path`` for reuse in a later session.
|
||||
|
||||
The file stores a plain dict with the audio tokens moved to CPU, so
|
||||
it can be loaded with ``torch.load(weights_only=True)`` (the default
|
||||
since torch 2.6) and is portable across devices.
|
||||
|
||||
Args:
|
||||
path: Destination file path (e.g. ``"my_voice.pt"``).
|
||||
"""
|
||||
torch.save(
|
||||
{
|
||||
"format_version": _VOICE_CLONE_PROMPT_FORMAT_VERSION,
|
||||
"ref_audio_tokens": self.ref_audio_tokens.detach().cpu(),
|
||||
"ref_text": self.ref_text,
|
||||
"ref_rms": float(self.ref_rms),
|
||||
},
|
||||
path,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def load(cls, path: str, map_location: str = "cpu") -> "VoiceClonePrompt":
|
||||
"""Load a prompt saved with :meth:`save`.
|
||||
|
||||
The returned prompt can be passed directly to
|
||||
:meth:`OmniVoice.generate`; the audio tokens are moved to the model
|
||||
device automatically during generation, so no manual ``.to(device)``
|
||||
is needed.
|
||||
|
||||
Args:
|
||||
path: File path previously written by :meth:`save`.
|
||||
map_location: Device to load the audio tokens onto.
|
||||
Returns:
|
||||
The restored :class:`VoiceClonePrompt`.
|
||||
"""
|
||||
data = torch.load(path, map_location=map_location, weights_only=True)
|
||||
version = data.get("format_version")
|
||||
if version != _VOICE_CLONE_PROMPT_FORMAT_VERSION:
|
||||
raise ValueError(f"Unsupported VoiceClonePrompt format version: {version}")
|
||||
return cls(
|
||||
ref_audio_tokens=data["ref_audio_tokens"],
|
||||
ref_text=data["ref_text"],
|
||||
ref_rms=data["ref_rms"],
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class OmniVoiceGenerationConfig:
|
||||
|
||||
@@ -0,0 +1,667 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FlashInfer-accelerated iterative decoding for OmniVoice.
|
||||
|
||||
Approach (mirrors CosyVoice/runtime/triton_trtllm/token2wav_dit_flashinfer.py):
|
||||
|
||||
- Sequence packing: the baseline pads the uncond (CFG) sequence to the cond
|
||||
length and runs batch=2 with a (2,1,S,S) bool mask. Here cond+uncond are
|
||||
packed into ONE row of length c_len+u_len with per-document positions and
|
||||
flashinfer ragged attention (qo_indptr = document boundaries) — no pad
|
||||
compute, no S^2 mask materialization.
|
||||
- Attention: registered as a custom HF attention implementation
|
||||
("omnivoice_fi") via AttentionInterface; reads the wrapper planned
|
||||
once per generation from a module-level context. HF mask construction is
|
||||
bypassed by passing attention_mask={"full_attention": None}.
|
||||
- KV cache: disabled (llm.config.use_cache=False). Iterative bidirectional
|
||||
decoding recomputes the full sequence every step, so the DynamicCache the
|
||||
baseline builds each forward is pure overhead.
|
||||
- Optional CUDA graphs: one graph per packed shape; all 32 denoising steps
|
||||
replay the same graph (input_ids/audio_mask/position_ids are copied into
|
||||
static buffers). Each shape owns a private flashinfer wrapper, since a
|
||||
plan bakes its launch metadata into the captured graph.
|
||||
|
||||
Usage:
|
||||
from omnivoice_flashinfer import apply_flashinfer
|
||||
apply_flashinfer(model, enable_cuda_graph=True)
|
||||
|
||||
Ported from upstream k2-fsa/OmniVoice master with one behavioural change:
|
||||
the unmasking schedule uses ``num_step + 1`` timesteps to match this repo's
|
||||
``_generate_iterative``. VoiceStudio enables it via ``OMNIVOICE_FLASHINFER``
|
||||
(see services/model_manager.py); ``flashinfer`` is an optional dependency and
|
||||
this module must only be imported after that opt-in.
|
||||
"""
|
||||
|
||||
import math
|
||||
import time
|
||||
from types import MethodType
|
||||
from typing import List
|
||||
|
||||
import flashinfer
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from transformers.modeling_utils import AttentionInterface
|
||||
|
||||
from omnivoice.models.omnivoice import (
|
||||
GenerationTask,
|
||||
OmniVoiceGenerationConfig,
|
||||
_get_time_steps,
|
||||
_gumbel_sample,
|
||||
)
|
||||
|
||||
_WORKSPACE_SIZE = 128 * 1024 * 1024
|
||||
# Context read by the registered attention function. "wrapper" must be planned
|
||||
# for the current packed layout before any llm forward.
|
||||
_CTX = {"wrapper": None}
|
||||
|
||||
|
||||
def _flashinfer_attention(
|
||||
module, query, key, value, attention_mask, scaling=None, dropout=0.0, **kwargs
|
||||
):
|
||||
"""query (1, Hq, S, D), key/value (1, Hkv, S, D) — packed documents."""
|
||||
_b, hq, s, d = query.shape
|
||||
hkv = key.shape[1]
|
||||
q = query.transpose(1, 2).reshape(s, hq, d)
|
||||
k = key.transpose(1, 2).reshape(s, hkv, d)
|
||||
v = value.transpose(1, 2).reshape(s, hkv, d)
|
||||
out = _CTX["wrapper"].run(q, k, v) # (S, Hq, D)
|
||||
return out.view(1, s, hq, d), None
|
||||
|
||||
|
||||
AttentionInterface.register("omnivoice_fi", _flashinfer_attention)
|
||||
|
||||
|
||||
def _fi_rmsnorm_forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
"""Single-kernel replacement for Qwen3RMSNorm.forward (a 7-kernel
|
||||
fp32-upcast chain in eager mode). flashinfer.norm.rmsnorm computes in
|
||||
fp32 internally and matches to fp16 rounding."""
|
||||
shape = hidden_states.shape
|
||||
out = flashinfer.norm.rmsnorm(
|
||||
hidden_states.reshape(-1, shape[-1]).contiguous(),
|
||||
self.weight,
|
||||
eps=self.variance_epsilon,
|
||||
)
|
||||
return out.view(shape)
|
||||
|
||||
|
||||
def _patch_rmsnorm(llm):
|
||||
from transformers.models.qwen3.modeling_qwen3 import Qwen3RMSNorm
|
||||
|
||||
n = 0
|
||||
for module in llm.modules():
|
||||
if isinstance(module, Qwen3RMSNorm):
|
||||
module.forward = MethodType(_fi_rmsnorm_forward, module)
|
||||
n += 1
|
||||
return n
|
||||
|
||||
|
||||
def _fi_attention_module_forward(
|
||||
self,
|
||||
hidden_states,
|
||||
position_embeddings=None,
|
||||
attention_mask=None,
|
||||
past_key_values=None,
|
||||
**kwargs,
|
||||
):
|
||||
"""NHD-layout replacement for Qwen3Attention.forward (packed batch=1).
|
||||
|
||||
The stock forward works in (B, H, S, D): the rotate-half RoPE costs a cat
|
||||
plus four elementwise passes, and handing (B,H,S,D) to the ragged wrapper
|
||||
costs three transpose copies. Keeping everything in (S, H, D) removes all
|
||||
of that; RoPE is one fused in-place kernel driven by packed position ids
|
||||
(read from _CTX, set per generation / baked per graph)."""
|
||||
s = hidden_states.shape[1]
|
||||
x = hidden_states[0] # (S, hidden)
|
||||
if getattr(self, "_fi_w_qkv", None) is not None:
|
||||
qkv = F.linear(x, self._fi_w_qkv)
|
||||
q, k, v = qkv.split(self._fi_qkv_split, dim=-1)
|
||||
# split views are strided; reshape materializes contiguous copies
|
||||
# (q/k would be copied inside the fused rmsnorm anyway)
|
||||
q = self.q_norm(q.reshape(s, -1, self.head_dim))
|
||||
k = self.k_norm(k.reshape(s, -1, self.head_dim))
|
||||
v = v.reshape(s, -1, self.head_dim)
|
||||
else:
|
||||
q = self.q_norm(self.q_proj(x).view(s, -1, self.head_dim))
|
||||
k = self.k_norm(self.k_proj(x).view(s, -1, self.head_dim))
|
||||
v = self.v_proj(x).view(s, -1, self.head_dim)
|
||||
flashinfer.rope.apply_rope_pos_ids_inplace(
|
||||
q, k, _CTX["pos_ids"], rope_theta=self._fi_rope_theta, interleave=False
|
||||
)
|
||||
slots = _CTX.get("doc_slots")
|
||||
if slots is not None:
|
||||
# bucketed-graph mode: a flashinfer plan bakes document boundaries
|
||||
# into the graph, so attention runs per fixed-length document slot as
|
||||
# SDPA with an O(slot) key-padding mask whose contents are rewritten
|
||||
# per generation. (A dense (S,S) block-diag mask scales quadratically
|
||||
# and the enable_gqa+mask combo drops SDPA to the math backend, so
|
||||
# k/v are pre-expanded to full heads instead.)
|
||||
ng = self.num_key_value_groups
|
||||
k = k.repeat_interleave(ng, dim=1) # (S, Hq, D)
|
||||
v = v.repeat_interleave(ng, dim=1)
|
||||
out = torch.empty_like(q)
|
||||
for start, slot_len, m in slots:
|
||||
od = F.scaled_dot_product_attention(
|
||||
q[start : start + slot_len].transpose(0, 1).unsqueeze(0),
|
||||
k[start : start + slot_len].transpose(0, 1).unsqueeze(0),
|
||||
v[start : start + slot_len].transpose(0, 1).unsqueeze(0),
|
||||
attn_mask=m,
|
||||
)
|
||||
out[start : start + slot_len] = od.squeeze(0).transpose(0, 1)
|
||||
else:
|
||||
out = _CTX["wrapper"].run(q, k, v) # (S, Hq, D)
|
||||
return self.o_proj(out.reshape(s, -1)).unsqueeze(0), None
|
||||
|
||||
|
||||
def _patch_attention_forward(llm, fuse_qkv=True):
|
||||
theta = llm.config.rope_parameters["rope_theta"]
|
||||
for layer in llm.layers:
|
||||
attn = layer.self_attn
|
||||
attn._fi_rope_theta = theta
|
||||
if fuse_qkv:
|
||||
attn._fi_w_qkv = torch.cat(
|
||||
[attn.q_proj.weight, attn.k_proj.weight, attn.v_proj.weight], dim=0
|
||||
)
|
||||
attn._fi_qkv_split = [
|
||||
attn.q_proj.weight.shape[0],
|
||||
attn.k_proj.weight.shape[0],
|
||||
attn.v_proj.weight.shape[0],
|
||||
]
|
||||
attn.forward = MethodType(_fi_attention_module_forward, attn)
|
||||
|
||||
|
||||
def _fi_mlp_forward(self, x):
|
||||
"""Qwen3MLP with fused gate+up GEMM and flashinfer silu_and_mul
|
||||
(2 GEMMs + silu + mul -> 1 GEMM + 1 fused kernel)."""
|
||||
y = F.linear(x[0], self._fi_w_gate_up) # (S, 2*inter)
|
||||
y = flashinfer.activation.silu_and_mul(y)
|
||||
return self.down_proj(y).unsqueeze(0)
|
||||
|
||||
|
||||
def _patch_mlp(llm):
|
||||
for layer in llm.layers:
|
||||
mlp = layer.mlp
|
||||
mlp._fi_w_gate_up = torch.cat([mlp.gate_proj.weight, mlp.up_proj.weight], dim=0)
|
||||
mlp.forward = MethodType(_fi_mlp_forward, mlp)
|
||||
|
||||
|
||||
class PackedAttnRunner:
|
||||
def __init__(
|
||||
self,
|
||||
num_qo_heads,
|
||||
num_kv_heads,
|
||||
head_dim,
|
||||
device,
|
||||
workspace_size=_WORKSPACE_SIZE,
|
||||
):
|
||||
self.num_qo_heads = num_qo_heads
|
||||
self.num_kv_heads = num_kv_heads
|
||||
self.head_dim = head_dim
|
||||
self.device = device
|
||||
self._workspace = torch.zeros(workspace_size, dtype=torch.uint8, device=device)
|
||||
self.wrapper = flashinfer.BatchPrefillWithRaggedKVCacheWrapper(
|
||||
self._workspace, "NHD"
|
||||
)
|
||||
self._planned_key = None
|
||||
|
||||
def plan(self, doc_lens: List[int], dtype: torch.dtype):
|
||||
key = (tuple(doc_lens), dtype)
|
||||
if key == self._planned_key:
|
||||
return
|
||||
indptr = torch.zeros(len(doc_lens) + 1, dtype=torch.int32, device=self.device)
|
||||
indptr[1:] = torch.cumsum(
|
||||
torch.tensor(doc_lens, dtype=torch.int32, device=self.device), dim=0
|
||||
)
|
||||
self.wrapper.plan(
|
||||
indptr,
|
||||
indptr,
|
||||
self.num_qo_heads,
|
||||
self.num_kv_heads,
|
||||
self.head_dim,
|
||||
causal=False,
|
||||
sm_scale=self.head_dim**-0.5,
|
||||
q_data_type=dtype,
|
||||
kv_data_type=dtype,
|
||||
)
|
||||
self._planned_key = key
|
||||
|
||||
|
||||
def _generate_iterative_packed(
|
||||
self, task: GenerationTask, gen_config: OmniVoiceGenerationConfig
|
||||
) -> List[torch.Tensor]:
|
||||
"""Packed-sequence rewrite of OmniVoice._generate_iterative.
|
||||
|
||||
Documents are packed as [cond_0, uncond_0, cond_1, uncond_1, ...] into a
|
||||
single batch row; the scoring/unmasking math is identical to the original.
|
||||
"""
|
||||
B = task.batch_size
|
||||
inputs_list = [
|
||||
self._prepare_inference_inputs(
|
||||
task.texts[i],
|
||||
task.target_lens[i],
|
||||
task.ref_texts[i],
|
||||
task.ref_audio_tokens[i],
|
||||
task.langs[i],
|
||||
task.instructs[i],
|
||||
gen_config.denoise,
|
||||
)
|
||||
for i in range(B)
|
||||
]
|
||||
|
||||
c_lens = [inp["input_ids"].size(2) for inp in inputs_list]
|
||||
u_lens = list(task.target_lens)
|
||||
doc_lens = []
|
||||
for c, u in zip(c_lens, u_lens):
|
||||
doc_lens.extend([c, u])
|
||||
|
||||
use_graph = getattr(self, "_fi_enable_cuda_graph", False)
|
||||
buckets = getattr(self, "_fi_graph_buckets", None) # durations in seconds
|
||||
|
||||
# Choose the packed layout. Bucketed-graph mode places each item in fixed
|
||||
# slots [C_budget | U_budget] so one graph per (batch, duration bucket)
|
||||
# serves any sample that fits; otherwise pack tightly.
|
||||
bucket_U = None
|
||||
if use_graph and buckets is not None:
|
||||
frame_rate = self.audio_tokenizer.config.frame_rate
|
||||
t_max = max(u_lens)
|
||||
overhead_max = max(c - u for c, u in zip(c_lens, u_lens))
|
||||
bucket_U = next(
|
||||
(int(d * frame_rate) for d in sorted(buckets) if d * frame_rate >= t_max),
|
||||
None,
|
||||
)
|
||||
if bucket_U is None or overhead_max > self._fi_overhead_budget:
|
||||
bucket_U = None
|
||||
use_graph = False # too long for the buckets: eager fallback
|
||||
|
||||
if bucket_U is not None:
|
||||
U_b = bucket_U
|
||||
C_b = U_b + self._fi_overhead_budget
|
||||
offsets = []
|
||||
for i in range(B):
|
||||
offsets.extend([i * (C_b + U_b), i * (C_b + U_b) + C_b])
|
||||
total_len = B * (C_b + U_b)
|
||||
else:
|
||||
offsets = [0]
|
||||
for l in doc_lens[:-1]:
|
||||
offsets.append(offsets[-1] + l)
|
||||
total_len = sum(doc_lens)
|
||||
|
||||
C = self.config.num_audio_codebook
|
||||
packed_ids = torch.full(
|
||||
(1, C, total_len),
|
||||
self.config.audio_mask_id,
|
||||
dtype=torch.long,
|
||||
device=self.device,
|
||||
)
|
||||
packed_audio_mask = torch.zeros(
|
||||
(1, total_len), dtype=torch.bool, device=self.device
|
||||
)
|
||||
position_ids = torch.zeros((1, total_len), dtype=torch.long, device=self.device)
|
||||
|
||||
for i, inp in enumerate(inputs_list):
|
||||
c_off, u_off = offsets[2 * i], offsets[2 * i + 1]
|
||||
c_len, u_len = c_lens[i], u_lens[i]
|
||||
packed_ids[0, :, c_off : c_off + c_len] = inp["input_ids"][0]
|
||||
packed_audio_mask[0, c_off : c_off + c_len] = inp["audio_mask"][0]
|
||||
position_ids[0, c_off : c_off + c_len] = torch.arange(c_len, device=self.device)
|
||||
# uncond doc = target region only
|
||||
packed_ids[0, :, u_off : u_off + u_len] = inp["input_ids"][0, :, -u_len:]
|
||||
packed_audio_mask[0, u_off : u_off + u_len] = inp["audio_mask"][0, -u_len:]
|
||||
position_ids[0, u_off : u_off + u_len] = torch.arange(u_len, device=self.device)
|
||||
|
||||
# num_step + 1 mirrors our _generate_iterative's schedule (a local
|
||||
# divergence from upstream k2-fsa): packed decoding must unmask on exactly
|
||||
# the same schedule as the eager path or outputs differ between the two.
|
||||
timesteps = _get_time_steps(
|
||||
t_start=0.0,
|
||||
t_end=1.0,
|
||||
num_step=gen_config.num_step + 1,
|
||||
t_shift=gen_config.t_shift,
|
||||
).tolist()
|
||||
schedules = []
|
||||
for t_len in task.target_lens:
|
||||
total_mask = t_len * C
|
||||
rem = total_mask
|
||||
sched = []
|
||||
for step in range(gen_config.num_step):
|
||||
num = (
|
||||
rem
|
||||
if step == gen_config.num_step - 1
|
||||
else min(
|
||||
math.ceil(total_mask * (timesteps[step + 1] - timesteps[step])), rem
|
||||
)
|
||||
)
|
||||
sched.append(int(num))
|
||||
rem -= int(num)
|
||||
schedules.append(sched)
|
||||
|
||||
layer_ids = torch.arange(C, device=self.device).view(1, -1, 1)
|
||||
|
||||
# gather indices of the logits-consuming positions, laid out as
|
||||
# [all cond-target blocks | all uncond blocks] so the guidance/scoring
|
||||
# math can run over every item in one batched pass. flat_spans[i] gives
|
||||
# the item's (start, len) within each half; in bucket mode items sit at a
|
||||
# fixed stride U_b with junk rows (pointing at position 0) in between.
|
||||
cond_ranges, uncond_ranges = [], []
|
||||
flat_spans = []
|
||||
for i in range(B):
|
||||
c_off, u_off = offsets[2 * i], offsets[2 * i + 1]
|
||||
c_len, t_len = c_lens[i], task.target_lens[i]
|
||||
if bucket_U is not None:
|
||||
flat_spans.append((U_b * i, t_len))
|
||||
cond_rows = torch.zeros(U_b, dtype=torch.long, device=self.device)
|
||||
cond_rows[:t_len] = torch.arange(
|
||||
c_off + c_len - t_len, c_off + c_len, device=self.device
|
||||
)
|
||||
uncond_rows = torch.zeros(U_b, dtype=torch.long, device=self.device)
|
||||
uncond_rows[:t_len] = torch.arange(u_off, u_off + t_len, device=self.device)
|
||||
cond_ranges.append(cond_rows)
|
||||
uncond_ranges.append(uncond_rows)
|
||||
else:
|
||||
prev = 0 if i == 0 else flat_spans[-1][0] + flat_spans[-1][1]
|
||||
flat_spans.append((prev, t_len))
|
||||
cond_ranges.append(
|
||||
torch.arange(c_off + c_len - t_len, c_off + c_len, device=self.device)
|
||||
)
|
||||
uncond_ranges.append(torch.arange(u_off, u_off + t_len, device=self.device))
|
||||
T_flat = (U_b * B) if bucket_U is not None else sum(task.target_lens)
|
||||
tgt_index = torch.cat(cond_ranges + uncond_ranges)
|
||||
|
||||
# flat per-position token state aligned with the cond half of the gathered
|
||||
# layout. Junk positions (bucket-mode slot padding) are initialized to -1
|
||||
# so the global "already unmasked" fill gives them -inf scores and topk
|
||||
# never selects them.
|
||||
tokens_flat = torch.full((C, T_flat), -1, dtype=torch.long, device=self.device)
|
||||
for st, t_len in flat_spans:
|
||||
tokens_flat[:, st : st + t_len] = self.config.audio_mask_id
|
||||
|
||||
if use_graph and bucket_U is not None:
|
||||
graph_entry = _get_or_capture_bucket_graph(self, B, U_b, C_b)
|
||||
# refresh the per-generation static contents (shape-invariant, data-variant)
|
||||
graph_entry["audio_mask"].copy_(packed_audio_mask)
|
||||
graph_entry["position_ids"].copy_(position_ids)
|
||||
graph_entry["pos_ids_i32"].copy_(position_ids[0].to(torch.int32))
|
||||
graph_entry["tgt_index"].copy_(tgt_index)
|
||||
for d_idx, m in enumerate(graph_entry["doc_masks"]):
|
||||
length = c_lens[d_idx // 2] if d_idx % 2 == 0 else u_lens[d_idx // 2]
|
||||
m[..., :length] = True
|
||||
m[..., length:] = False
|
||||
elif use_graph:
|
||||
graph_entry = _get_or_capture_graph(self, tuple(doc_lens), tgt_index)
|
||||
graph_entry["audio_mask"].copy_(packed_audio_mask)
|
||||
graph_entry["position_ids"].copy_(position_ids)
|
||||
else:
|
||||
self._fi_runner.plan(doc_lens, torch.float16)
|
||||
_CTX["wrapper"] = self._fi_runner.wrapper
|
||||
_CTX["pos_ids"] = position_ids[0].to(torch.int32)
|
||||
_CTX["doc_slots"] = None
|
||||
|
||||
# optional llm timing hook (set by the benchmark; graph replays bypass
|
||||
# model.forward, so wrapping forward would miss them)
|
||||
stats = getattr(self, "_fi_llm_stats", None)
|
||||
|
||||
for step in range(gen_config.num_step):
|
||||
if stats is not None:
|
||||
torch.cuda.synchronize()
|
||||
t0 = time.perf_counter()
|
||||
if use_graph:
|
||||
graph_entry["input_ids"].copy_(packed_ids)
|
||||
graph_entry["graph"].replay()
|
||||
batch_logits = graph_entry["logits"].to(torch.float32)
|
||||
else:
|
||||
batch_logits = _forward_logits(
|
||||
self, packed_ids, packed_audio_mask, position_ids, tgt_index
|
||||
).to(torch.float32)
|
||||
if stats is not None:
|
||||
torch.cuda.synchronize()
|
||||
stats["seconds"] += time.perf_counter() - t0
|
||||
stats["calls"] += 1
|
||||
|
||||
# batched scoring over every item at once: the guidance/log_softmax/
|
||||
# argmax/gumbel chain (the GPU-heavy part) runs on the whole
|
||||
# [cond | uncond] halves; only topk + scatter stay per item.
|
||||
c_logits_all = batch_logits[:, :, :T_flat, :]
|
||||
u_logits_all = batch_logits[:, :, T_flat:, :]
|
||||
pred_all, scores_all = self._predict_tokens_with_scoring(
|
||||
c_logits_all, u_logits_all, gen_config
|
||||
)
|
||||
scores_all = scores_all - (layer_ids * gen_config.layer_penalty_factor)
|
||||
if gen_config.position_temperature > 0.0:
|
||||
scores_all = _gumbel_sample(scores_all, gen_config.position_temperature)
|
||||
# -inf for already-unmasked positions AND bucket-slot junk (-1)
|
||||
scores_all.masked_fill_(
|
||||
(tokens_flat != self.config.audio_mask_id).unsqueeze(0), -float("inf")
|
||||
)
|
||||
pred_all, scores_all = pred_all[0], scores_all[0] # (C, T_flat)
|
||||
|
||||
for i in range(B):
|
||||
k = schedules[i][step]
|
||||
if k <= 0:
|
||||
continue
|
||||
c_off, u_off = offsets[2 * i], offsets[2 * i + 1]
|
||||
c_len, t_len = c_lens[i], task.target_lens[i]
|
||||
st, _ = flat_spans[i]
|
||||
|
||||
_, topk_idx = torch.topk(scores_all[:, st : st + t_len].reshape(-1), k)
|
||||
flat_tokens = tokens_flat[:, st : st + t_len].reshape(-1)
|
||||
flat_tokens[topk_idx] = pred_all[:, st : st + t_len].reshape(-1)[topk_idx]
|
||||
new_tokens = flat_tokens.view(C, t_len)
|
||||
tokens_flat[:, st : st + t_len] = new_tokens
|
||||
|
||||
packed_ids[0, :, c_off + c_len - t_len : c_off + c_len] = new_tokens
|
||||
packed_ids[0, :, u_off : u_off + t_len] = new_tokens
|
||||
|
||||
return [tokens_flat[:, st : st + t_len] for (st, t_len) in flat_spans]
|
||||
|
||||
|
||||
def _forward_logits(model, input_ids, audio_mask, position_ids, tgt_index):
|
||||
"""LLM forward + audio head over target positions only.
|
||||
|
||||
The scoring step consumes logits at the cond-target and uncond ranges
|
||||
(2*sum(t_len) of the packed positions); running the 1024->8200 audio_heads
|
||||
GEMM and the fp32 upcast on the full packed length is wasted work.
|
||||
Returns logits of shape (1, C, 2*sum(t_len), V) laid out as
|
||||
[all cond-target blocks | all uncond blocks] — matching tgt_index
|
||||
(torch.cat(cond_ranges + uncond_ranges)) and the caller's split at T_flat.
|
||||
"""
|
||||
inputs_embeds = model._prepare_embed_inputs(input_ids, audio_mask)
|
||||
hidden = model.llm(
|
||||
inputs_embeds=inputs_embeds,
|
||||
attention_mask={"full_attention": None},
|
||||
return_dict=True,
|
||||
position_ids=position_ids,
|
||||
)[0]
|
||||
tgt_hidden = hidden[0, tgt_index] # (2T, hidden)
|
||||
logits_flat = model.audio_heads(tgt_hidden)
|
||||
n = tgt_hidden.shape[0]
|
||||
return logits_flat.view(
|
||||
1, n, model.config.num_audio_codebook, model.config.audio_vocab_size
|
||||
).permute(0, 2, 1, 3)
|
||||
|
||||
|
||||
def _get_or_capture_graph(model, doc_lens_key, tgt_index):
|
||||
cache = model._fi_graph_cache
|
||||
entry = cache.get(doc_lens_key)
|
||||
if entry is not None:
|
||||
return entry
|
||||
|
||||
device = model.device
|
||||
total_len = sum(doc_lens_key)
|
||||
C = model.config.num_audio_codebook
|
||||
llm_cfg = model.config.llm_config
|
||||
runner = PackedAttnRunner(
|
||||
llm_cfg.num_attention_heads,
|
||||
llm_cfg.num_key_value_heads,
|
||||
llm_cfg.head_dim,
|
||||
device,
|
||||
workspace_size=64 * 1024 * 1024,
|
||||
)
|
||||
runner.plan(list(doc_lens_key), torch.float16)
|
||||
_CTX["wrapper"] = runner.wrapper
|
||||
|
||||
# positions are fully determined by doc_lens (the cache key), so both the
|
||||
# long buffer (model-level rotary) and the int32 copy (fused rope) can be
|
||||
# baked with their final values
|
||||
positions = torch.cat([torch.arange(l, device=device) for l in doc_lens_key])
|
||||
static = {
|
||||
"input_ids": torch.full(
|
||||
(1, C, total_len),
|
||||
model.config.audio_mask_id,
|
||||
dtype=torch.long,
|
||||
device=device,
|
||||
),
|
||||
"audio_mask": torch.zeros((1, total_len), dtype=torch.bool, device=device),
|
||||
"position_ids": positions.unsqueeze(0).contiguous(),
|
||||
}
|
||||
pos_ids_i32 = positions.to(torch.int32)
|
||||
_CTX["pos_ids"] = pos_ids_i32
|
||||
_CTX["doc_slots"] = None
|
||||
|
||||
side_stream = torch.cuda.Stream()
|
||||
side_stream.wait_stream(torch.cuda.current_stream())
|
||||
with torch.cuda.stream(side_stream):
|
||||
for _ in range(2):
|
||||
_forward_logits(
|
||||
model,
|
||||
static["input_ids"],
|
||||
static["audio_mask"],
|
||||
static["position_ids"],
|
||||
tgt_index,
|
||||
)
|
||||
torch.cuda.current_stream().wait_stream(side_stream)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
# tgt_index depends only on doc_lens (the cache key), so it is safe
|
||||
# to bake into the graph
|
||||
logits = _forward_logits(
|
||||
model,
|
||||
static["input_ids"],
|
||||
static["audio_mask"],
|
||||
static["position_ids"],
|
||||
tgt_index,
|
||||
)
|
||||
|
||||
# tgt_index is baked into the captured gather by pointer — the entry must
|
||||
# keep it alive or the allocator will reuse its memory for later samples
|
||||
entry = {
|
||||
"graph": graph,
|
||||
"logits": logits,
|
||||
"runner": runner,
|
||||
"tgt_index": tgt_index,
|
||||
"pos_ids_i32": pos_ids_i32,
|
||||
**static,
|
||||
}
|
||||
cache[doc_lens_key] = entry
|
||||
return entry
|
||||
|
||||
|
||||
def _get_or_capture_bucket_graph(model, B, U_b, C_b):
|
||||
"""One graph per (batch, duration-bucket): items sit in fixed
|
||||
[C_budget | U_budget] slots; attention runs as SDPA over a runtime-updated
|
||||
block-diagonal mask, so any sample that fits the slots replays exactly."""
|
||||
key = ("bucket", B, U_b)
|
||||
cache = model._fi_graph_cache
|
||||
entry = cache.get(key)
|
||||
if entry is not None:
|
||||
return entry
|
||||
|
||||
device = model.device
|
||||
total_len = B * (C_b + U_b)
|
||||
C = model.config.num_audio_codebook
|
||||
|
||||
static = {
|
||||
"input_ids": torch.full(
|
||||
(1, C, total_len),
|
||||
model.config.audio_mask_id,
|
||||
dtype=torch.long,
|
||||
device=device,
|
||||
),
|
||||
"audio_mask": torch.zeros((1, total_len), dtype=torch.bool, device=device),
|
||||
"position_ids": torch.zeros((1, total_len), dtype=torch.long, device=device),
|
||||
"pos_ids_i32": torch.zeros(total_len, dtype=torch.int32, device=device),
|
||||
"tgt_index": torch.zeros(2 * B * U_b, dtype=torch.long, device=device),
|
||||
}
|
||||
# per-document key-padding masks (contents updated per generation);
|
||||
# init all-True so warmup/capture has no fully-masked softmax rows
|
||||
doc_masks, doc_slots = [], []
|
||||
for i in range(B):
|
||||
for slot_start, slot_len in (
|
||||
(i * (C_b + U_b), C_b),
|
||||
(i * (C_b + U_b) + C_b, U_b),
|
||||
):
|
||||
m = torch.ones(1, 1, 1, slot_len, dtype=torch.bool, device=device)
|
||||
doc_masks.append(m)
|
||||
doc_slots.append((slot_start, slot_len, m))
|
||||
|
||||
_CTX["wrapper"] = None
|
||||
_CTX["pos_ids"] = static["pos_ids_i32"]
|
||||
_CTX["doc_slots"] = doc_slots
|
||||
|
||||
side_stream = torch.cuda.Stream()
|
||||
side_stream.wait_stream(torch.cuda.current_stream())
|
||||
with torch.cuda.stream(side_stream):
|
||||
for _ in range(2):
|
||||
_forward_logits(
|
||||
model,
|
||||
static["input_ids"],
|
||||
static["audio_mask"],
|
||||
static["position_ids"],
|
||||
static["tgt_index"],
|
||||
)
|
||||
torch.cuda.current_stream().wait_stream(side_stream)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
logits = _forward_logits(
|
||||
model,
|
||||
static["input_ids"],
|
||||
static["audio_mask"],
|
||||
static["position_ids"],
|
||||
static["tgt_index"],
|
||||
)
|
||||
|
||||
entry = {
|
||||
"graph": graph,
|
||||
"logits": logits,
|
||||
"doc_masks": doc_masks,
|
||||
"doc_slots": doc_slots,
|
||||
**static,
|
||||
}
|
||||
cache[key] = entry
|
||||
return entry
|
||||
|
||||
|
||||
def apply_flashinfer(
|
||||
model,
|
||||
enable_cuda_graph: bool = False,
|
||||
fuse_rmsnorm: bool = True,
|
||||
fuse_attention: bool = True,
|
||||
cuda_graph_buckets=None,
|
||||
overhead_budget: int = 512,
|
||||
):
|
||||
"""Patch an OmniVoice instance to use flashinfer packed attention."""
|
||||
model.llm.set_attn_implementation("omnivoice_fi")
|
||||
if fuse_rmsnorm:
|
||||
_patch_rmsnorm(model.llm)
|
||||
if fuse_attention:
|
||||
_patch_attention_forward(model.llm)
|
||||
_patch_mlp(model.llm)
|
||||
# Bidirectional iterative decoding recomputes everything each step; the
|
||||
# DynamicCache the baseline allocates+fills per forward is pure overhead.
|
||||
model.llm.config.use_cache = False
|
||||
|
||||
llm_cfg = model.config.llm_config
|
||||
model._fi_runner = PackedAttnRunner(
|
||||
llm_cfg.num_attention_heads,
|
||||
llm_cfg.num_key_value_heads,
|
||||
llm_cfg.head_dim,
|
||||
model.device,
|
||||
)
|
||||
model._fi_graph_cache = {}
|
||||
model._fi_enable_cuda_graph = enable_cuda_graph or cuda_graph_buckets is not None
|
||||
model._fi_graph_buckets = cuda_graph_buckets
|
||||
model._fi_overhead_budget = overhead_budget
|
||||
model._generate_iterative = MethodType(_generate_iterative_packed, model)
|
||||
return model
|
||||
@@ -0,0 +1,214 @@
|
||||
"""The FlashInfer opt-in (OMNIVOICE_FLASHINFER, upstream k2-fsa port).
|
||||
|
||||
An optimization must never be a point of failure (#278 contract, same as
|
||||
torch.compile): the env knob is CUDA-only, off by default, refuses with a
|
||||
named reason when the host can't honor it, latches off for the session after
|
||||
a runtime failure, and a mid-generation FlashInfer error unapplies the patch
|
||||
and retries the standard path once.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
torch = pytest.importorskip("torch")
|
||||
|
||||
|
||||
def _ee():
|
||||
import services.engine_env as m
|
||||
return m
|
||||
|
||||
|
||||
def _mm():
|
||||
import services.model_manager as m
|
||||
return m
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_latch(monkeypatch):
|
||||
monkeypatch.setattr(_ee(), "_flashinfer_runtime_failure", None)
|
||||
monkeypatch.delenv("OMNIVOICE_FLASHINFER", raising=False)
|
||||
|
||||
|
||||
# ── the env knob ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"value,expected",
|
||||
[
|
||||
("", "off"), ("0", "off"), ("false", "off"), ("off", "off"),
|
||||
("1", "on"), ("true", "on"), ("ON", "on"),
|
||||
("graph", "graph"), ("GRAPH", "graph"),
|
||||
("banana", "off"), # typo → default path, not a crash
|
||||
],
|
||||
)
|
||||
def test_flashinfer_mode_parsing(monkeypatch, value, expected):
|
||||
if value:
|
||||
monkeypatch.setenv("OMNIVOICE_FLASHINFER", value)
|
||||
assert _ee().flashinfer_mode() == expected
|
||||
|
||||
|
||||
def test_should_flashinfer_refuses_non_cuda(monkeypatch):
|
||||
monkeypatch.setenv("OMNIVOICE_FLASHINFER", "1")
|
||||
assert _ee().should_flashinfer("cpu") == "off"
|
||||
assert _ee().should_flashinfer("mps") == "off"
|
||||
|
||||
|
||||
def test_should_flashinfer_refuses_without_the_package(monkeypatch):
|
||||
monkeypatch.setenv("OMNIVOICE_FLASHINFER", "1")
|
||||
ee = _ee()
|
||||
monkeypatch.setattr(ee.importlib.util, "find_spec", lambda name: None)
|
||||
assert ee.should_flashinfer("cuda") == "off"
|
||||
|
||||
|
||||
def test_latched_reason_is_sanitized(monkeypatch):
|
||||
# Wheel import errors embed the user's home path — the latch must store
|
||||
# the redacted form (core.failure.sanitize maps $HOME → "~").
|
||||
import os
|
||||
|
||||
home = os.path.expanduser("~")
|
||||
_ee().mark_flashinfer_runtime_failure(
|
||||
f"ImportError: {home}/.venv/lib/flashinfer/_kernels.so: bad ELF"
|
||||
)
|
||||
latched = _ee()._flashinfer_runtime_failure
|
||||
assert home not in latched
|
||||
assert "ImportError" in latched
|
||||
|
||||
|
||||
def test_sanitizer_failure_never_latches_the_raw_reason(monkeypatch):
|
||||
# Fail closed: a broken redactor must not leak the original message.
|
||||
import core.failure
|
||||
|
||||
def _boom(_):
|
||||
raise RuntimeError("sanitizer exploded (test)")
|
||||
|
||||
monkeypatch.setattr(core.failure, "sanitize", _boom)
|
||||
_ee().mark_flashinfer_runtime_failure(
|
||||
"ImportError: /home/someone/secret-project/creds.so missing"
|
||||
)
|
||||
latched = _ee()._flashinfer_runtime_failure
|
||||
assert "secret-project" not in latched and "/home/" not in latched
|
||||
assert latched.startswith("ImportError")
|
||||
assert "redacted" in latched
|
||||
|
||||
|
||||
def test_runtime_failure_latches_the_session_off(monkeypatch):
|
||||
monkeypatch.setenv("OMNIVOICE_FLASHINFER", "graph")
|
||||
ee = _ee()
|
||||
monkeypatch.setattr(ee.importlib.util, "find_spec", lambda name: object())
|
||||
assert ee.should_flashinfer("cuda") == "graph"
|
||||
ee.mark_flashinfer_runtime_failure("boom")
|
||||
assert ee.should_flashinfer("cuda") == "off"
|
||||
|
||||
|
||||
# ── failure classification ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_classifier_matches_flashinfer_markers():
|
||||
mm = _mm()
|
||||
assert mm._is_flashinfer_runtime_failure(RuntimeError("flashinfer plan failed"))
|
||||
assert mm._is_flashinfer_runtime_failure(RuntimeError("CUDA graph capture aborted"))
|
||||
assert not mm._is_flashinfer_runtime_failure(ValueError("Unsupported instruct items"))
|
||||
assert not mm._is_flashinfer_runtime_failure(RuntimeError("CUDA out of memory"))
|
||||
|
||||
|
||||
def test_classifier_walks_the_cause_chain():
|
||||
mm = _mm()
|
||||
inner = RuntimeError("flashinfer workspace too small")
|
||||
outer = RuntimeError("generation failed")
|
||||
outer.__cause__ = inner
|
||||
assert mm._is_flashinfer_runtime_failure(outer)
|
||||
# `raise ... from None` severs the chain — a genuine error must not be
|
||||
# re-classified via a suppressed FlashInfer context.
|
||||
severed = RuntimeError("generation failed")
|
||||
severed.__context__ = inner
|
||||
severed.__suppress_context__ = True
|
||||
assert not mm._is_flashinfer_runtime_failure(severed)
|
||||
|
||||
|
||||
# ── unapply restores the class implementations ──────────────────────────────
|
||||
|
||||
|
||||
class _MiniModel:
|
||||
class _Llm(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.lin = torch.nn.Linear(2, 2)
|
||||
self.config = type("C", (), {"use_cache": False})()
|
||||
self.attn_impl = None
|
||||
|
||||
def set_attn_implementation(self, name):
|
||||
self.attn_impl = name
|
||||
|
||||
def __init__(self):
|
||||
self.llm = self._Llm()
|
||||
|
||||
def _generate_iterative(self, *a):
|
||||
return "class-impl"
|
||||
|
||||
|
||||
def test_unapply_flashinfer_restores_instance_state():
|
||||
from types import MethodType
|
||||
|
||||
m = _MiniModel()
|
||||
# Simulate apply_flashinfer's instance-level patching.
|
||||
m.llm.lin.forward = MethodType(lambda self, x: "patched", m.llm.lin)
|
||||
m.llm.lin._fi_w_qkv = torch.zeros(1)
|
||||
m._generate_iterative = MethodType(lambda self, *a: "patched", m)
|
||||
m._fi_runner = object()
|
||||
m._fi_graph_cache = {}
|
||||
m._fi_enable_cuda_graph = True
|
||||
|
||||
_mm()._unapply_flashinfer(m)
|
||||
|
||||
assert "forward" not in vars(m.llm.lin), "instance forward override must go"
|
||||
assert not hasattr(m.llm.lin, "_fi_w_qkv")
|
||||
assert m._generate_iterative() == "class-impl"
|
||||
assert not hasattr(m, "_fi_runner")
|
||||
assert m.llm.attn_impl == "sdpa"
|
||||
assert m.llm.config.use_cache is True
|
||||
|
||||
|
||||
def test_unapply_restores_the_captured_attention_impl():
|
||||
# The pre-apply impl may be flash_attention_2, not sdpa — unapply must
|
||||
# put back what was actually there (CodeRabbit/Greptile, #1565).
|
||||
m = _MiniModel()
|
||||
m._fi_orig_attn_impl = "flash_attention_2"
|
||||
_mm()._unapply_flashinfer(m)
|
||||
assert m.llm.attn_impl == "flash_attention_2"
|
||||
assert not hasattr(m, "_fi_orig_attn_impl")
|
||||
|
||||
|
||||
# ── generate-time fallback ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_generate_fallback_unapplies_and_retries_once():
|
||||
mm = _mm()
|
||||
calls = {"n": 0}
|
||||
|
||||
class _Model(_MiniModel):
|
||||
def generate(self, **kw):
|
||||
calls["n"] += 1
|
||||
if calls["n"] == 1:
|
||||
raise RuntimeError("flashinfer ragged attention failed")
|
||||
return ["ok"]
|
||||
|
||||
m = _Model()
|
||||
m._fi_runner = object()
|
||||
mm._install_flashinfer_fallback(m)
|
||||
assert m.generate() == ["ok"]
|
||||
assert calls["n"] == 2
|
||||
assert not hasattr(m, "_fi_runner"), "fallback must unapply the patch"
|
||||
assert _ee()._flashinfer_runtime_failure is not None
|
||||
|
||||
|
||||
def test_generate_fallback_leaves_real_errors_alone():
|
||||
mm = _mm()
|
||||
|
||||
class _Model(_MiniModel):
|
||||
def generate(self, **kw):
|
||||
raise ValueError("Unsupported instruct items")
|
||||
|
||||
m = _Model()
|
||||
mm._install_flashinfer_fallback(m)
|
||||
with pytest.raises(ValueError):
|
||||
m.generate()
|
||||
@@ -50,6 +50,7 @@ _ALLOWED_FILES = {
|
||||
"README_CN.md", # Chinese README (a translation)
|
||||
"docs/data_preparation.md", # multilingual example payloads
|
||||
"docs/voice-design.md", # EN/CJK attribute mapping table
|
||||
"docs/engines/omnivoice.md", # pinyin pronunciation-control example (functional CJK)
|
||||
"docs/superpowers/specs/2026-05-31-voice-gallery-design.md", # Chinese-dialect taxonomy reference table
|
||||
"examples/README.md", # multilingual example payloads
|
||||
# Text-processing (CJK punctuation inside sentence/clause-splitting regexes)
|
||||
|
||||
@@ -0,0 +1,223 @@
|
||||
"""Voice-clone prompts persist across restarts (upstream VoiceClonePrompt port).
|
||||
|
||||
The in-memory prompt cache (#427/#473) dies with the process, so the first
|
||||
generation of every session re-encoded each voice — and re-ran ASR when the
|
||||
profile had no stored transcript. Upstream k2-fsa added
|
||||
``VoiceClonePrompt.save()/.load()`` for exactly this; we port the format
|
||||
(version-tagged dict, ``torch.load(weights_only=True)``-safe) and put a disk
|
||||
layer under the memory LRU, keyed identically (ref path + mtime + ref_text +
|
||||
preprocess flag). Restart is simulated here by clearing the memory cache: a
|
||||
second lookup must come from disk, not a re-encode.
|
||||
|
||||
The layer is best-effort by contract: disabled (env), unwritable, or corrupt
|
||||
disk state must never fail a generation — worst case is the old re-encode.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
torch = pytest.importorskip("torch")
|
||||
|
||||
|
||||
def _tb():
|
||||
"""The *live* services.tts_backend (same rationale as
|
||||
test_clone_prompt_wiring._tb: other suites purge services.* modules)."""
|
||||
import services.tts_backend as m
|
||||
return m
|
||||
|
||||
|
||||
def _VoiceClonePrompt():
|
||||
"""Resolved at call time — a module-level binding could go stale when
|
||||
another suite purges omnivoice.* from sys.modules (CodeRabbit, #1565)."""
|
||||
from omnivoice.models.omnivoice import VoiceClonePrompt
|
||||
return VoiceClonePrompt
|
||||
|
||||
|
||||
def _prompt():
|
||||
return _VoiceClonePrompt()(
|
||||
ref_audio_tokens=torch.arange(24, dtype=torch.long).reshape(8, 3),
|
||||
ref_text="Nice to meet you.",
|
||||
ref_rms=0.123,
|
||||
)
|
||||
|
||||
|
||||
class _StubModel:
|
||||
def __init__(self):
|
||||
self.encodes = 0
|
||||
|
||||
def create_voice_clone_prompt(self, ref_audio, ref_text=None, preprocess_prompt=True):
|
||||
self.encodes += 1
|
||||
return _prompt()
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _isolated(tmp_path, monkeypatch):
|
||||
"""Point the disk layer at a per-test dir and start with empty caches."""
|
||||
monkeypatch.setattr("core.config.DATA_DIR", tmp_path / "data")
|
||||
monkeypatch.delenv("OMNIVOICE_PROMPT_DISK_CACHE", raising=False)
|
||||
_tb().clear_clone_prompt_cache()
|
||||
yield
|
||||
_tb().clear_clone_prompt_cache()
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def ref_wav(tmp_path):
|
||||
p = tmp_path / "ref.wav"
|
||||
p.write_bytes(b"\x00" * 256)
|
||||
return str(p)
|
||||
|
||||
|
||||
def _disk_files(tmp_path):
|
||||
d = tmp_path / "data" / "prompt_cache"
|
||||
return sorted(d.glob("*.pt")) if d.is_dir() else []
|
||||
|
||||
|
||||
# ── the ported save/load format ─────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_prompt_save_load_roundtrip(tmp_path):
|
||||
p = _prompt()
|
||||
path = str(tmp_path / "voice.pt")
|
||||
p.save(path)
|
||||
loaded = _VoiceClonePrompt().load(path)
|
||||
assert torch.equal(loaded.ref_audio_tokens, p.ref_audio_tokens)
|
||||
assert loaded.ref_text == p.ref_text
|
||||
assert loaded.ref_rms == pytest.approx(p.ref_rms)
|
||||
# The file must stay loadable under torch's safe default (weights_only=True
|
||||
# since 2.6) — a pickled dataclass would not be.
|
||||
raw = torch.load(path, weights_only=True)
|
||||
assert raw["format_version"] == 1
|
||||
|
||||
|
||||
def test_prompt_load_rejects_unknown_format_version(tmp_path):
|
||||
path = str(tmp_path / "future.pt")
|
||||
torch.save({"format_version": 999}, path)
|
||||
with pytest.raises(ValueError, match="format version"):
|
||||
_VoiceClonePrompt().load(path)
|
||||
|
||||
|
||||
def test_saved_tokens_are_cpu_even_from_dataclass_on_another_device(tmp_path):
|
||||
# save() must detach+CPU the tokens so the file is portable. On CUDA hosts
|
||||
# this exercises the real device move; CI (CPU-only) still verifies the
|
||||
# detach and that the persisted payload is CPU-resident.
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
p = _VoiceClonePrompt()(
|
||||
ref_audio_tokens=torch.zeros(8, 3, requires_grad=True).to(device),
|
||||
ref_text="x",
|
||||
ref_rms=0.5,
|
||||
)
|
||||
path = str(tmp_path / "v.pt")
|
||||
p.save(path)
|
||||
loaded = _VoiceClonePrompt().load(path)
|
||||
assert not loaded.ref_audio_tokens.requires_grad
|
||||
assert loaded.ref_audio_tokens.device.type == "cpu"
|
||||
# The device move must happen at SAVE time (portability of the file
|
||||
# itself), not merely at load: the raw payload carries CPU tensors.
|
||||
assert torch.load(path, weights_only=True)["ref_audio_tokens"].device.type == "cpu"
|
||||
|
||||
|
||||
# ── the disk layer under the memory cache ───────────────────────────────────
|
||||
|
||||
|
||||
def test_disk_hit_survives_restart(tmp_path, ref_wav):
|
||||
tb = _tb()
|
||||
model = _StubModel()
|
||||
|
||||
first = tb._get_clone_prompt(model, ref_wav, "hello", True)
|
||||
assert model.encodes == 1
|
||||
assert len(_disk_files(tmp_path)) == 1
|
||||
|
||||
tb.clear_clone_prompt_cache() # "restart": memory gone, disk remains
|
||||
second = tb._get_clone_prompt(model, ref_wav, "hello", True)
|
||||
assert model.encodes == 1, "restart re-encoded despite a persisted prompt"
|
||||
assert torch.equal(second.ref_audio_tokens, first.ref_audio_tokens)
|
||||
assert second.ref_text == first.ref_text
|
||||
|
||||
|
||||
def test_edited_reference_is_not_served_a_stale_prompt(tmp_path, ref_wav):
|
||||
import os
|
||||
|
||||
tb = _tb()
|
||||
model = _StubModel()
|
||||
tb._get_clone_prompt(model, ref_wav, "hello", True)
|
||||
tb.clear_clone_prompt_cache()
|
||||
|
||||
# Same path, new content+mtime → new key → the old file must not match.
|
||||
with open(ref_wav, "wb") as f:
|
||||
f.write(b"\x01" * 512)
|
||||
os.utime(ref_wav, (1, 1))
|
||||
tb._get_clone_prompt(model, ref_wav, "hello", True)
|
||||
assert model.encodes == 2
|
||||
|
||||
|
||||
def test_single_use_refs_never_touch_disk(tmp_path, ref_wav):
|
||||
tb = _tb()
|
||||
tb._get_clone_prompt(_StubModel(), ref_wav, "hello", True, store=False)
|
||||
assert _disk_files(tmp_path) == [], (
|
||||
"store=False (dub per-segment clips) must not spray single-use "
|
||||
"prompts onto disk — same scan-resistance as the memory LRU"
|
||||
)
|
||||
|
||||
|
||||
def test_env_kill_switch_disables_the_layer(tmp_path, ref_wav, monkeypatch):
|
||||
monkeypatch.setenv("OMNIVOICE_PROMPT_DISK_CACHE", "0")
|
||||
tb = _tb()
|
||||
model = _StubModel()
|
||||
tb._get_clone_prompt(model, ref_wav, "hello", True)
|
||||
assert _disk_files(tmp_path) == []
|
||||
tb.clear_clone_prompt_cache()
|
||||
tb._get_clone_prompt(model, ref_wav, "hello", True)
|
||||
assert model.encodes == 2 # no disk → honest re-encode
|
||||
|
||||
|
||||
def test_corrupt_disk_entry_is_dropped_and_reencoded(tmp_path, ref_wav):
|
||||
tb = _tb()
|
||||
model = _StubModel()
|
||||
tb._get_clone_prompt(model, ref_wav, "hello", True)
|
||||
tb.clear_clone_prompt_cache()
|
||||
|
||||
disk = _disk_files(tmp_path)
|
||||
assert len(disk) == 1
|
||||
disk[0].write_bytes(b"not a torch file")
|
||||
|
||||
prompt = tb._get_clone_prompt(model, ref_wav, "hello", True)
|
||||
assert prompt is not None
|
||||
assert model.encodes == 2, "corrupt file must fall back to encoding"
|
||||
# ...and the corrupt file was removed, then replaced by the fresh save.
|
||||
fresh = _disk_files(tmp_path)
|
||||
assert len(fresh) == 1
|
||||
assert torch.load(str(fresh[0]), weights_only=True)["format_version"] == 1
|
||||
|
||||
|
||||
def test_prune_keeps_only_the_newest(tmp_path, monkeypatch):
|
||||
import os
|
||||
import time
|
||||
|
||||
tb = _tb()
|
||||
monkeypatch.setattr(tb, "_PROMPT_DISK_CACHE_MAX", 3)
|
||||
model = _StubModel()
|
||||
refs = []
|
||||
for i in range(5):
|
||||
p = tmp_path / f"ref{i}.wav"
|
||||
p.write_bytes(bytes([i]) * 64)
|
||||
os.utime(p, (i + 1, i + 1))
|
||||
refs.append(str(p))
|
||||
for i, r in enumerate(refs):
|
||||
tb._get_clone_prompt(model, r, f"text {i}", True)
|
||||
# mtime is the prune order; keep saves strictly ordered even on
|
||||
# filesystems with coarse timestamps.
|
||||
files = _disk_files(tmp_path)
|
||||
newest = max(files, key=lambda f: f.stat().st_mtime)
|
||||
os.utime(newest, (1000 + i, 1000 + i))
|
||||
assert len(_disk_files(tmp_path)) == 3
|
||||
|
||||
|
||||
def test_unwritable_cache_dir_never_breaks_prompt_building(ref_wav, monkeypatch):
|
||||
# Simulate an unwritable data dir: the layer must vanish, not raise.
|
||||
monkeypatch.setattr(
|
||||
"core.config.DATA_DIR", "/proc/omnivoice-definitely-not-writable"
|
||||
)
|
||||
tb = _tb()
|
||||
model = _StubModel()
|
||||
assert tb._get_clone_prompt(model, ref_wav, "hello", True) is not None
|
||||
assert model.encodes == 1
|
||||
Reference in New Issue
Block a user