feat(chatterbox-fast): Phase 1 streaming server — adaptive-chunk scheduler

Build the streaming TTS server MVP per docs/design/chatterbox-fast-plan.md §4.

- scheduler.py: adaptive buffer-ratchet chunker (the meat) — GPU-free pure
  logic. First sentence emitted alone for low TTFA, then chunks ratchet ~3x by
  packing whole sentences to margin x buffered-audio; drives off measured RTF +
  sec/char (EMA). relieve_leader() clause-splits a too-big mid-stream sentence
  to avoid starvation (joins land on commas); a long comma-less sentence is the
  one honored-but-flagged limitation.
- test_scheduler.py: GPU-free simulation, 13 tests — asserts no-starvation
  (incl. overestimated RTF) and the ratchet.
- app.py: FastAPI model holder + POST /tts StreamingResponse (raw PCM s16le
  default, wav optional, stream/oneshot) + GET /health.
- bench.py: client — ground-truth TTFB + real 1x-consumer starvation check.

Live test on irv-ml1 (turbo, A6000, GLaDOS voice): streaming TTFB 499ms vs
oneshot 5230ms (~10x), stayed ahead of a 1x player (no starvation), ratchet
1.64->4.08->8.60->8.60s audio, measured RTF self-corrected 3.38->4.01.

Kill the superseded docs/design/chatterbox-fast.md — its §5 windowed-token
streaming was the abandoned native-frame-streaming arc; the adaptive-chunk plan
supersedes it. Repoint persistent-memory + README at the canonical plan.
This commit is contained in:
vh
2026-06-01 22:42:28 -07:00
parent 4b9bd109bf
commit 7cd39001b2
7 changed files with 957 additions and 141 deletions
+96
View File
@@ -0,0 +1,96 @@
# chatterbox-fast — streaming TTS engine
Custom streaming server on top of `ChatterboxTurboTTS` that delivers
**sub-second time-to-first-audio** while keeping turbo's full quality. Workload:
**single-stream interactive**. Deployed (Phase 3) **alongside** the live
`chatterbox` (:8196) on irv-ml1, burned in, then catalog-flipped.
Design: [`docs/design/chatterbox-fast-plan.md`](../../docs/design/chatterbox-fast-plan.md)
(canonical plan). The abandoned native-frame-streaming arc is recorded in
`persistent-memory.md` → *Tried and abandoned*.
## How it works — adaptive buffer-ratchet chunking
The engine never splits mid-sentence (keeps each chunk prosodically coherent).
Instead it rides Chatterbox's faster-than-realtime generation (RTF ~3.4–3.8×):
1. **Chunk 1 = first sentence**, generated and emitted immediately (~0.66s
first-audio). Latency-critical.
2. **While chunk N plays, generate chunk N+1** by greedily accumulating whole
sentences until the next would exceed `margin × audio_buffered_remaining`.
3. Each chunk's playback buys wall-clock for a ~3× bigger next chunk, so after
2-3 joins the rest of the paragraph is one big near-full-context chunk.
Context loss is confined to those few sentence-boundary joins.
4. Driven off **measured** RTF + sec-per-char (EMA), not constants.
5. **Starvation relief:** if a mid-stream sentence is too long to generate
within the current buffer, its *clause* boundaries are exposed so chunks pack
to commas (natural pauses) — never a mid-clause split. A long *comma-less*
sentence after a short opener is the one unavoidable case: the rule is
honored and the brief gap is **flagged** (`drained > 0`), never hidden.
This only works because RTF > 1 — a sub-realtime model (e.g. Fish) would starve
regardless of chunking. That is why this is the chatterbox-specific answer.
## Files
| file | role |
|---|---|
| `scheduler.py` | The adaptive-chunk scheduler. **GPU-free, pure logic** — the meat. |
| `test_scheduler.py` | GPU-free simulation: asserts no-starvation + ratchet. `python test_scheduler.py` or `pytest`. |
| `app.py` | FastAPI server: model holder + `POST /tts` (StreamingResponse) + `GET /health`. |
| `bench.py` | Client: ground-truth TTFB + real 1×-consumer starvation check; saves `.wav` for A/B. |
Phase 3 will add `compose.yaml`, `Dockerfile`, `.env.example`.
## API
`POST /tts` → streamed audio. Body:
```json
{ "text": "...", "voice": "glados_25s", "format": "pcm",
"stream": true, "exaggeration": 0.5, "temperature": 0.8,
"top_p": 0.95, "top_k": 1000, "repetition_penalty": 1.2 }
```
- `format`: `pcm` (raw s16le @ 24 kHz, lowest latency, default) or `wav`.
- `stream: false` → whole-text one-shot (the A/B quality baseline).
- `voice`: predefined name (a `*.wav` in `CBF_VOICES_DIR`) or an absolute path
to a clone reference. Omit → server default.
- `margin` / `margin_first` / `rtf_prior`: optional scheduler overrides.
`GET /health` → `{status, sr, device, default_voice, voices_dir}`.
## Config (env)
| var | default | meaning |
|---|---|---|
| `CBF_MODEL_DEVICE` | `cuda` | `cuda` / `cuda:0` / `cpu` |
| `CBF_VOICES_DIR` | `/refs` | dir of predefined voice wavs |
| `CBF_DEFAULT_VOICE` | first wav in dir | default reference wav (path or name) |
| `CBF_BIND` / `CBF_PORT` | `0.0.0.0` / `8197` | uvicorn bind |
## Dev / test on irv-ml1
```bash
# (from this dir) copy the server into the chatterbox image and run it on GPU 1:
scp app.py scheduler.py bench.py lkraven@10.100.79.3:/tmp/cbf/
IMG=$(ssh lkraven@10.100.79.3 "docker images --format '{{.Repository}}:{{.Tag}}' | grep -i chatterbox | grep -v '<none>' | head -1")
ssh lkraven@10.100.79.3 "docker run --rm --gpus '\"device=1\"' -e NVIDIA_VISIBLE_DEVICES=1 \
-e HF_HOME=/app/hf_cache -e CBF_VOICES_DIR=/refs -e CBF_DEFAULT_VOICE=glados_25s \
-p 8197:8197 -v /worktank/chatterbox/cache:/app/hf_cache \
-v /worktank/chatterbox/reference_audio:/refs -v /tmp/cbf:/cbf \
$IMG python /cbf/app.py"
# then, from the host (or anywhere on the WG net):
python bench.py --host http://10.100.79.3:8197 --out /refs/_fast.wav
python bench.py --host http://10.100.79.3:8197 --oneshot --out /refs/_oneshot.wav
```
Pull the samples to listen: `scp lkraven@10.100.79.3:/worktank/chatterbox/reference_audio/_*.wav ~/chatterbox-ab/`.
## Acceptance (plan §6)
- **Latency:** first-audio < ~0.8s on the deployment GPU.
- **No starvation:** `bench.py` reports "stayed ahead"; `test_scheduler.py` green.
- **Quality:** operator ear-A/B the streamed output vs the one-shot — join-context
loss should be ~imperceptible for multi-sentence text.
+270
View File
@@ -0,0 +1,270 @@
"""chatterbox-fast — streaming TTS server (Phase 1 MVP).
A lean FastAPI server on the ChatterboxTurboTTS library that streams audio using
the adaptive buffer-ratchet scheduler in ``scheduler.py`` (the meat). Sub-second
time-to-first-audio while keeping turbo's full quality; workload is single-stream
interactive (see docs/design/chatterbox-fast-plan.md).
Endpoints:
POST /tts — StreamingResponse of audio chunks (raw PCM s16le default).
GET /health — model/voice readiness.
Config via env (all optional; sane dev defaults):
CBF_MODEL_DEVICE cuda | cuda:0 | cpu (default: cuda)
CBF_VOICES_DIR dir of predefined voice wavs (default: /refs)
CBF_DEFAULT_VOICE default reference wav path/name (default: first wav in dir)
CBF_BIND / CBF_PORT uvicorn bind (default: 0.0.0.0:8197)
"""
from __future__ import annotations
import logging
import os
import struct
import threading
import time
from contextlib import asynccontextmanager
from pathlib import Path
from typing import Iterator, Literal
import torch
from fastapi import FastAPI, HTTPException
from fastapi.responses import StreamingResponse
from pydantic import BaseModel, Field
from scheduler import ChunkConfig, ChunkResult, stream_chunks
log = logging.getLogger("chatterbox-fast")
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
# ── config ────────────────────────────────────────────────────────────────
DEVICE = os.environ.get("CBF_MODEL_DEVICE", "cuda")
VOICES_DIR = Path(os.environ.get("CBF_VOICES_DIR", "/refs"))
DEFAULT_VOICE_ENV = os.environ.get("CBF_DEFAULT_VOICE")
BIND = os.environ.get("CBF_BIND", "0.0.0.0")
PORT = int(os.environ.get("CBF_PORT", "8197"))
# Turbo sampling knobs validated in the spike (plan §2). CFG / exaggeration /
# min_p are ignored by turbo (it warns, harmless).
WARMUP_TEXT = "Warming up the streaming engine."
# ── model holder ──────────────────────────────────────────────────────────
class Engine:
"""Owns the single ChatterboxTurboTTS instance + a generation lock.
The workload is single-stream interactive, but prepare_conditionals mutates
model state, so all generation is serialized under one lock. Concurrent
multi-voice serving is a Phase-2 concern.
"""
def __init__(self) -> None:
self.model = None
self.sr = 24000
self.lock = threading.Lock()
self._current_voice: str | None = None
self.default_voice: str | None = None
def load(self) -> None:
from chatterbox.tts_turbo import ChatterboxTurboTTS
log.info("loading ChatterboxTurboTTS on %s …", DEVICE)
t0 = time.perf_counter()
self.model = ChatterboxTurboTTS.from_pretrained(device=DEVICE)
self.sr = int(getattr(self.model, "sr", 24000))
self.default_voice = self._discover_default_voice()
if self.default_voice:
self._prepare(self.default_voice)
self._warm()
log.info(
"model ready in %.1fs (sr=%d, default_voice=%s)",
time.perf_counter() - t0, self.sr, self.default_voice,
)
def _discover_default_voice(self) -> str | None:
if DEFAULT_VOICE_ENV:
# Resolve a bare name ("glados_25s") the same way request-time voices
# are resolved — append .wav and look under VOICES_DIR.
return self.resolve_voice(DEFAULT_VOICE_ENV)
if VOICES_DIR.is_dir():
wavs = sorted(VOICES_DIR.glob("*.wav"))
if wavs:
return str(wavs[0])
return None
def resolve_voice(self, voice: str | None) -> str:
if not voice:
if not self.default_voice:
raise HTTPException(503, "no default voice; set CBF_DEFAULT_VOICE")
return self.default_voice
p = Path(voice)
if p.is_absolute() and p.exists():
return str(p)
cand = VOICES_DIR / (voice if voice.endswith(".wav") else f"{voice}.wav")
if cand.exists():
return str(cand)
raise HTTPException(404, f"unknown voice {voice!r}")
def _prepare(self, voice_path: str, exaggeration: float = 0.5) -> None:
if voice_path == self._current_voice:
return
log.info("prepare_conditionals(%s)", voice_path)
self.model.prepare_conditionals(voice_path, exaggeration=exaggeration, norm_loudness=True)
self._current_voice = voice_path
def _warm(self) -> None:
with torch.inference_mode():
self.model.generate(WARMUP_TEXT, repetition_penalty=1.2, top_p=0.95,
temperature=0.8, top_k=1000)
if DEVICE.startswith("cuda"):
torch.cuda.synchronize()
def generate(self, text: str, knobs: "TTSRequest") -> tuple[torch.Tensor, float]:
"""Synthesize ``text`` → (wav tensor [1,T], audio_seconds). CUDA-synced
so the caller's clock delta is honest gen time."""
with torch.inference_mode():
wav = self.model.generate(
text,
repetition_penalty=knobs.repetition_penalty,
top_p=knobs.top_p,
temperature=knobs.temperature,
top_k=knobs.top_k,
)
if DEVICE.startswith("cuda"):
torch.cuda.synchronize()
audio_sec = wav.shape[-1] / self.sr
return wav, audio_sec
engine = Engine()
@asynccontextmanager
async def lifespan(app: FastAPI):
engine.load()
yield
app = FastAPI(title="chatterbox-fast", lifespan=lifespan)
# ── request / audio encoding ──────────────────────────────────────────────
class TTSRequest(BaseModel):
text: str
voice: str | None = None
format: Literal["pcm", "wav"] = "pcm"
stream: bool = True # False ⇒ whole-text one-shot (for A/B vs streaming)
exaggeration: float = 0.5
temperature: float = 0.8
top_p: float = 0.95
top_k: int = 1000
repetition_penalty: float = 1.2
# Scheduler overrides (None ⇒ ChunkConfig defaults).
margin: float | None = Field(default=None)
margin_first: float | None = Field(default=None)
rtf_prior: float | None = Field(default=None)
def _pcm16(wav: torch.Tensor) -> bytes:
a = wav.detach().to(torch.float32).clamp_(-1.0, 1.0).cpu().numpy().reshape(-1)
return (a * 32767.0).astype("<i2").tobytes()
def _wav_header(sr: int) -> bytes:
"""Streaming WAV header with unknown length (0xFFFFFFFF sizes)."""
return (
b"RIFF" + struct.pack("<I", 0xFFFFFFFF) + b"WAVE"
+ b"fmt " + struct.pack("<IHHIIHH", 16, 1, 1, sr, sr * 2, 2, 16)
+ b"data" + struct.pack("<I", 0xFFFFFFFF)
)
def _chunk_config(req: TTSRequest) -> ChunkConfig:
cfg = ChunkConfig()
if req.margin is not None:
cfg.margin = req.margin
if req.margin_first is not None:
cfg.margin_first = req.margin_first
if req.rtf_prior is not None:
cfg.rtf_prior = req.rtf_prior
return cfg
# ── endpoints ─────────────────────────────────────────────────────────────
@app.get("/health")
def health() -> dict:
return {
"status": "ok" if engine.model is not None else "loading",
"sr": engine.sr,
"device": DEVICE,
"default_voice": engine.default_voice,
"voices_dir": str(VOICES_DIR),
}
@app.post("/tts")
def tts(req: TTSRequest) -> StreamingResponse:
if engine.model is None:
raise HTTPException(503, "model still loading")
if not req.text.strip():
raise HTTPException(400, "empty text")
voice_path = engine.resolve_voice(req.voice)
cfg = _chunk_config(req)
media = "audio/wav" if req.format == "wav" else "application/octet-stream"
def body() -> Iterator[bytes]:
# One request holds the lock for its whole stream (single-stream
# workload); concurrent callers queue rather than corrupt conditionals.
with engine.lock:
engine._prepare(voice_path, exaggeration=req.exaggeration)
if req.format == "wav":
yield _wav_header(engine.sr)
t_req = time.perf_counter()
first_audio_ms: float | None = None
if not req.stream:
wav, audio_sec = engine.generate(req.text, req)
first_audio_ms = (time.perf_counter() - t_req) * 1000
log.info("oneshot: %.0fms gen, %.2fs audio", first_audio_ms, audio_sec)
yield _pcm16(wav)
return
def _gen(text: str) -> tuple[torch.Tensor, float]:
return engine.generate(text, req)
total_audio = 0.0
for r in stream_chunks(req.text, generate=_gen, clock=time.perf_counter, cfg=cfg):
if first_audio_ms is None:
first_audio_ms = (time.perf_counter() - t_req) * 1000
total_audio += r.audio_sec
_log_chunk(r, first_audio_ms)
yield _pcm16(r.audio)
log.info("stream done: ttfa=%.0fms total_audio=%.2fs", first_audio_ms or 0, total_audio)
return StreamingResponse(body(), media_type=media)
def _log_chunk(r: ChunkResult, ttfa_ms: float) -> None:
tag = " STARVED" if r.starved else ""
if r.index == 0:
log.info("chunk 0: ttfa=%.0fms gen=%.0fms audio=%.2fs rtf=%.2f%s",
ttfa_ms, r.gen_time * 1000, r.audio_sec, r.rtf, tag)
else:
log.info("chunk %d: gen=%.0fms audio=%.2fs buf=%.2f→%.2f rtf=%.2f%s",
r.index, r.gen_time * 1000, r.audio_sec,
r.buffer_before, r.buffer_after, r.rtf, tag)
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host=BIND, port=PORT)
+106
View File
@@ -0,0 +1,106 @@
"""Client-side bench for chatterbox-fast — ground-truth TTFB + starvation check.
The server's scheduler reasons about starvation from an *estimated* buffer; this
client measures the real thing: time-to-first-byte over the wire, and whether a
true 1×-realtime consumer ever runs dry. Saves the streamed audio to a .wav so
the operator can ear-A/B it against the one-shot.
Usage:
python bench.py --host http://10.100.79.3:8197 --text "..." --out /refs/_fast.wav
python bench.py --host ... --oneshot --out /refs/_oneshot.wav # A/B baseline
Stdlib only (urllib + wave) so it runs anywhere, incl. inside the container.
"""
from __future__ import annotations
import argparse
import json
import time
import urllib.request
import wave
SR = 24000
BYTES_PER_SEC = SR * 2 # s16le mono
DEFAULT_TEXT = (
"The cake is a lie. "
"I am being entirely sincere, without a trace of sarcasm, when I say that "
"this is the single most important scientific breakthrough in the entire "
"history of this facility. "
"You will be baked, and then there will be cake. "
"It is delicious and moist, assuming you survive the testing protocol, which "
"the available data suggests you almost certainly will not."
)
def run(host: str, text: str, out: str, *, oneshot: bool, voice: str | None) -> None:
payload = {"text": text, "format": "pcm", "stream": not oneshot}
if voice:
payload["voice"] = voice
req = urllib.request.Request(
f"{host}/tts",
data=json.dumps(payload).encode(),
headers={"Content-Type": "application/json"},
method="POST",
)
t0 = time.perf_counter()
t_first = t0
ttfb: float | None = None
received = 0 # bytes
worst_lead = float("inf") # min (buffered_audio_s - elapsed_since_first_s)
starve_events = 0
pcm = bytearray()
with urllib.request.urlopen(req) as resp:
while True:
buf = resp.read(4096)
if not buf:
break
now = time.perf_counter()
if ttfb is None:
ttfb = now - t0
t_first = now
received += len(buf)
pcm += buf
buffered_s = received / BYTES_PER_SEC
elapsed_s = now - t_first
lead = buffered_s - elapsed_s # >0 ⇒ ahead of a 1× player
if lead < worst_lead:
worst_lead = lead
if lead < 0:
starve_events += 1
total = time.perf_counter() - t0
audio_s = received / BYTES_PER_SEC
if ttfb is None:
raise SystemExit("no audio received from server")
with wave.open(out, "wb") as w:
w.setnchannels(1)
w.setsampwidth(2)
w.setframerate(SR)
w.writeframes(bytes(pcm))
mode = "oneshot" if oneshot else "stream"
print(f"[{mode}] ttfb={ttfb*1000:.0f}ms audio={audio_s:.2f}s "
f"wall={total:.2f}s rtf={audio_s/total:.2f}x")
if not oneshot:
verdict = "OK (stayed ahead)" if starve_events == 0 else f"STARVED ({starve_events} reads dry)"
print(f" worst lead over 1x player = {worst_lead:.2f}s → {verdict}")
print(f" wrote {out}")
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--host", default="http://10.100.79.3:8197")
ap.add_argument("--text", default=DEFAULT_TEXT)
ap.add_argument("--out", default="/refs/_fast.wav")
ap.add_argument("--voice", default=None)
ap.add_argument("--oneshot", action="store_true", help="whole-text one-shot baseline")
args = ap.parse_args()
run(args.host, args.text, args.out, oneshot=args.oneshot, voice=args.voice)
if __name__ == "__main__":
main()
+277
View File
@@ -0,0 +1,277 @@
"""Adaptive buffer-ratchet chunk scheduler for chatterbox-fast.
This module is the *meat* of the streaming engine and is deliberately free of
any GPU / torch / chatterbox / FastAPI imports so the no-starvation guarantee
can be validated as a pure simulation (see ``test_scheduler.py``).
The design (docs/design/chatterbox-fast-plan.md §1):
* Chunk 1 = the first sentence alone — generated and emitted immediately so
first-audio latency is minimal.
* While chunk N plays, generate chunk N+1 by greedily accumulating WHOLE
sentences until the next sentence would push estimated gen-time past
``margin × audio_buffered_remaining``. Never split mid-sentence — each
chunk stays prosodically self-coherent and joins land at natural pauses.
* Because Chatterbox runs faster than realtime (RTF > 1), every chunk's
playback buys wall-clock for a larger next chunk, so chunks ratchet up
~3× and after 2-3 joins the rest of the paragraph is one big near-full-
context chunk. Context loss is confined to those 2-3 sentence-boundary
joins.
* Drive off the *measured* realtime factor and sec-per-char, not constants —
track them live and self-correct.
The scheduler is *online*: chunk N+1's boundary depends on the RTF measured
while generating chunk N, so it cannot be precomputed. ``stream_chunks`` runs
the loop with ``generate`` and ``clock`` injected, which is what makes the
simulation possible.
"""
from __future__ import annotations
import re
from dataclasses import dataclass
from typing import Callable, Iterator, Sequence
# ── tunables ──────────────────────────────────────────────────────────────
@dataclass
class ChunkConfig:
"""Scheduler tunables. Priors are conservative; live measurement corrects."""
# Realtime factor prior. 3090 ~3.4×, A6000 ~3.8× (plan §2). Start
# conservative — under-estimating RTF makes chunks smaller and safer.
rtf_prior: float = 3.4
# Safety fraction of the buffer to spend on the next chunk's generation.
# margin < 1 leaves headroom so estimate error doesn't starve the stream.
margin: float = 0.8
# Tighter margin on the FIRST transition (chunk 1 → chunk 2): the buffer is
# smallest there, so starvation risk is highest (plan §1.5: ~0.6-0.7).
margin_first: float = 0.65
# Seconds of audio per character of text. ~0.060 ≈ 16.6 chars/sec speech.
# Calibrated live per request.
sec_per_char_prior: float = 0.060
# EMA weight on the newest measurement when updating rtf / sec_per_char.
ema_alpha: float = 0.4
# If the first sentence's estimated audio exceeds this, clause-split it to
# protect first-audio latency (the ONLY place we split below sentence
# granularity — plan §1.1).
max_first_sec: float = 2.0
# ── result record ─────────────────────────────────────────────────────────
@dataclass
class ChunkResult:
"""Telemetry + payload for one emitted chunk."""
index: int
text: str
audio: object # opaque payload from generate() (wav tensor, fake, …)
audio_sec: float
gen_time: float
est_gen: float # gen-time the scheduler predicted before generating
buffer_before: float # unplayed audio (s) when this chunk's gen started
buffer_after: float # unplayed audio (s) once this chunk is emitted
drained: float # seconds the buffer ran dry during gen (>0 ⇒ starvation)
rtf: float # measured RTF after this chunk
sec_per_char: float # measured sec/char after this chunk
@property
def starved(self) -> bool:
return self.drained > 1e-9
# ── text splitting ────────────────────────────────────────────────────────
# Split after sentence-final punctuation when the next non-space looks like a
# new sentence start: a capital/digit (optionally behind an open quote) or an
# inline tag like "[laugh]". Pragmatic, not perfect: abbreviations ("Mr.",
# "e.g.") can over-split — a Phase-2 watch-out, harmless to quality (just an
# extra natural-pause join).
_SENTENCE_END = re.compile(r'(?<=[.!?])["\')\]]*\s+(?=["\'(]*(?:[A-Z0-9]|\[))')
# Clause boundaries for the first-sentence latency fallback only.
_CLAUSE_END = re.compile(r'(?<=[,;:])\s+')
def split_sentences(text: str) -> list[str]:
"""Split into sentence units, preserving punctuation. Whitespace-collapsed."""
text = " ".join(text.split())
if not text:
return []
return [s for s in (p.strip() for p in _SENTENCE_END.split(text)) if s]
def _split_clauses(sentence: str) -> list[str]:
return [c for c in (p.strip() for p in _CLAUSE_END.split(sentence)) if c]
def protect_first_audio(units: list[str], cfg: ChunkConfig) -> list[str]:
"""Clause-split the first unit if it's too long to hit the first-audio target.
Only the *leading* unit is split; the remainder is untouched. If the first
sentence has no clause boundary we accept the latency rather than split
mid-clause (quality > latency once we're past the budget).
"""
if not units:
return units
if len(units[0]) * cfg.sec_per_char_prior <= cfg.max_first_sec:
return units
pieces = _split_clauses(units[0])
if len(pieces) <= 1:
return units # nothing to split on; keep the long first sentence whole
return pieces + units[1:]
# ── chunk planning ────────────────────────────────────────────────────────
def _est_gen_time(text: str, *, rtf: float, sec_per_char: float) -> float:
return (len(text) * sec_per_char) / rtf
def plan_chunk(
remaining: Sequence[str],
buffer_remaining: float,
*,
margin: float,
rtf: float,
sec_per_char: float,
) -> tuple[str, list[str]]:
"""Greedily accumulate whole units until the next would blow the budget.
Always returns at least one unit (never empty, never splits a unit). With
``buffer_remaining == 0`` (the first chunk) the budget is 0, so exactly the
first unit is taken — which is the latency-critical chunk-1 rule.
"""
budget = margin * buffer_remaining
chunk = [remaining[0]]
i = 1
while i < len(remaining):
candidate = " ".join(chunk + [remaining[i]])
if _est_gen_time(candidate, rtf=rtf, sec_per_char=sec_per_char) > budget:
break
chunk.append(remaining[i])
i += 1
return " ".join(chunk), list(remaining[i:])
def relieve_leader(
remaining: list[str],
buffer_remaining: float,
*,
rtf: float,
sec_per_char: float,
) -> list[str]:
"""Expose a too-big leading sentence's clause boundaries to avoid starvation.
If generating ``remaining[0]`` alone would overrun the buffer (an audible
gap), clause-split it in place so ``plan_chunk`` can pack clause-pieces up to
the buffer — joins then land on commas (natural pauses) instead of producing
a gap. If the sentence has no clause boundary it's returned unchanged: we do
NOT split mid-clause (design rule), and the caller accepts + flags the rare
starvation in telemetry.
"""
leader = remaining[0]
if _est_gen_time(leader, rtf=rtf, sec_per_char=sec_per_char) <= buffer_remaining:
return remaining
pieces = _split_clauses(leader)
if len(pieces) <= 1:
return remaining
return pieces + remaining[1:]
def _ema(old: float, new: float, alpha: float) -> float:
return (1 - alpha) * old + alpha * new
# ── the online loop ───────────────────────────────────────────────────────
# generate(text) -> (audio_payload, audio_seconds)
GenerateFn = Callable[[str], "tuple[object, float]"]
ClockFn = Callable[[], float]
def stream_chunks(
text: str,
*,
generate: GenerateFn,
clock: ClockFn,
cfg: ChunkConfig | None = None,
) -> Iterator[ChunkResult]:
"""Run the adaptive buffer-ratchet schedule, yielding one ChunkResult per chunk.
``generate`` does the actual synthesis (real GPU or a simulation) and
returns its opaque audio payload plus the audio's duration in seconds.
``clock`` returns monotonically increasing seconds; ``gen_time`` is measured
as the clock delta around ``generate``.
Buffer model (all in audio-seconds):
* Playback begins when chunk 1 arrives, so no drain occurs during chunk 1.
* For chunk N≥2, the client plays ``gen_time`` seconds of buffered audio
while we generate it, then we add this chunk's audio. If ``gen_time``
exceeds the buffer the stream starved (``drained`` > 0).
"""
cfg = cfg or ChunkConfig()
units = protect_first_audio(split_sentences(text), cfg)
rtf = cfg.rtf_prior
sec_per_char = cfg.sec_per_char_prior
buffer_remaining = 0.0
remaining: list[str] = units
index = 0
while remaining:
first = index == 0
if not first:
# If the next indivisible sentence would starve the buffer, expose
# its clause boundaries so plan_chunk can pack to the buffer.
remaining = relieve_leader(
remaining, buffer_remaining, rtf=rtf, sec_per_char=sec_per_char
)
margin = cfg.margin_first if first else cfg.margin
chunk_text, remaining = plan_chunk(
remaining, buffer_remaining, margin=margin, rtf=rtf, sec_per_char=sec_per_char
)
est_gen = _est_gen_time(chunk_text, rtf=rtf, sec_per_char=sec_per_char)
t0 = clock()
audio, audio_sec = generate(chunk_text)
gen_time = clock() - t0
# Starvation: did the buffer run dry while we generated this chunk?
drained = 0.0 if first else max(0.0, gen_time - buffer_remaining)
# Live self-correction.
if gen_time > 0:
rtf = _ema(rtf, audio_sec / gen_time, cfg.ema_alpha)
if chunk_text:
sec_per_char = _ema(sec_per_char, audio_sec / len(chunk_text), cfg.ema_alpha)
buffer_before = buffer_remaining
if first:
buffer_remaining = audio_sec
else:
buffer_remaining = max(0.0, buffer_remaining - gen_time) + audio_sec
yield ChunkResult(
index=index,
text=chunk_text,
audio=audio,
audio_sec=audio_sec,
gen_time=gen_time,
est_gen=est_gen,
buffer_before=buffer_before,
buffer_after=buffer_remaining,
drained=drained,
rtf=rtf,
sec_per_char=sec_per_char,
)
index += 1
+202
View File
@@ -0,0 +1,202 @@
"""GPU-free simulation of the adaptive-chunk scheduler.
Validates the two acceptance properties from the plan (§6) without a GPU:
* NO STARVATION — the stream stays ahead of 1× playback.
* RATCHET — chunks grow after the latency-critical first one.
Run directly (``python test_scheduler.py``) or under pytest.
"""
from __future__ import annotations
from scheduler import (
ChunkConfig,
plan_chunk,
protect_first_audio,
split_sentences,
stream_chunks,
)
# A multi-sentence paragraph with varied lengths. The long sentences carry
# commas (as natural prose does), so the clause-split starvation relief has
# boundaries to work with.
PARAGRAPH = (
"The cake is a lie. "
"I am being entirely sincere, without a trace of sarcasm, when I say "
"that this is the single most important scientific breakthrough in the "
"entire history of this facility. "
"You will be baked, and then there will be cake. "
"It is delicious and moist, assuming you survive the testing protocol, "
"which the available data suggests you almost certainly will not. "
"Goodbye."
)
class FakeClock:
"""A clock the fake generator advances, so gen_time reflects simulated work."""
def __init__(self) -> None:
self.t = 0.0
def __call__(self) -> float:
return self.t
def make_generator(clock: FakeClock, *, true_rtf: float, sec_per_char: float):
"""A fake generate() that costs realistic wall-clock and returns audio_sec.
Audio duration is proportional to text length; generation costs
``audio_sec / true_rtf`` of (simulated) wall-clock, advancing the clock.
"""
def generate(text: str):
audio_sec = len(text) * sec_per_char
clock.t += audio_sec / true_rtf
return None, audio_sec
return generate
def run(true_rtf: float = 3.8, sec_per_char: float = 0.060, cfg: ChunkConfig | None = None):
clock = FakeClock()
gen = make_generator(clock, true_rtf=true_rtf, sec_per_char=sec_per_char)
return list(stream_chunks(PARAGRAPH, generate=gen, clock=clock, cfg=cfg))
# ── splitting ─────────────────────────────────────────────────────────────
def test_split_sentences_count():
units = split_sentences(PARAGRAPH)
assert units[0] == "The cake is a lie."
assert units[-1] == "Goodbye."
assert len(units) == 5
def test_split_preserves_inline_tags():
units = split_sentences("Hello there. [laugh] That was funny. Bye.")
# The tag must not be torn from its sentence nor split on.
assert any("[laugh]" in u for u in units)
assert len(units) == 3
def test_protect_first_audio_clause_splits_long_opener():
cfg = ChunkConfig(max_first_sec=0.5) # force the guard
long_open = ["This is a long opener, with a clause, and another clause.", "Short."]
out = protect_first_audio(long_open, cfg)
assert len(out) > len(long_open)
assert out[0] == "This is a long opener,"
def test_protect_first_audio_keeps_unsplittable_opener_whole():
cfg = ChunkConfig(max_first_sec=0.1)
out = protect_first_audio(["No clause boundaries here at all."], cfg)
assert out == ["No clause boundaries here at all."]
# ── chunk planning ────────────────────────────────────────────────────────
def test_first_chunk_is_single_unit():
# Zero buffer ⇒ exactly the first unit, regardless of margin/rtf.
units = split_sentences(PARAGRAPH)
chunk, rest = plan_chunk(units, 0.0, margin=0.65, rtf=3.8, sec_per_char=0.06)
assert chunk == units[0]
assert len(rest) == len(units) - 1
def test_plan_never_returns_empty():
chunk, rest = plan_chunk(["Only one."], 0.0, margin=0.8, rtf=3.8, sec_per_char=0.06)
assert chunk == "Only one."
assert rest == []
# ── the acceptance properties ─────────────────────────────────────────────
def test_no_starvation_nominal():
"""Stream never starves when the RTF prior matches reality."""
results = run(true_rtf=3.8)
assert all(not r.starved for r in results), [
(r.index, r.drained) for r in results if r.starved
]
def test_no_starvation_when_rtf_overestimated():
"""Prior says 3.8× but the GPU only delivers 3.0× — live correction + margin
must still keep the stream fed. This is the property that matters: the
scheduler must be robust to an optimistic prior, not just a matching one."""
results = run(true_rtf=3.0)
assert all(not r.starved for r in results), [
(r.index, r.drained) for r in results if r.starved
]
def test_no_starvation_slow_gpu():
"""Even at 2.5× (well below prior) the margin absorbs it for this text."""
results = run(true_rtf=2.5)
assert all(not r.starved for r in results), [
(r.index, r.drained) for r in results if r.starved
]
def test_chunks_ratchet_up():
"""After the latency-critical first chunk, chunks grow (buffer-ratchet)."""
results = run(true_rtf=3.8)
assert len(results) >= 3, "paragraph should not collapse to one chunk"
# Chunk 1 is the smallest (single sentence); the second chunk is larger.
assert len(results[1].text) > len(results[0].text)
# The paragraph consolidates: the last chunk carries multiple sentences.
assert results[-1].audio_sec >= results[0].audio_sec
def test_first_chunk_low_latency():
"""First chunk is one short sentence ⇒ smallest gen time ⇒ fast first audio."""
results = run(true_rtf=3.8)
assert results[0].text == "The cake is a lie."
# In the sim, chunk 0's gen_time IS the time-to-first-audio, and it must be
# the smallest of all chunks (everything after it is ≥ a clause).
assert results[0].gen_time == min(r.gen_time for r in results)
assert results[0].gen_time < 0.4
def test_unsplittable_long_sentence_flags_starvation():
"""KNOWN LIMITATION (surfaced, not hidden): a long COMMA-LESS sentence right
after a short opener cannot be clause-split, so we honor the never-split-mid-
sentence rule and accept a brief gap — which MUST show up as drained>0 in
telemetry so it is measurable, never silent."""
clock = FakeClock()
gen = make_generator(clock, true_rtf=3.8, sec_per_char=0.060)
text = (
"Hi. "
"I am now going to speak one extremely long sentence with no clause "
"boundaries at all so that nothing in here can ever be split apart by "
"the scheduler no matter how hard it tries to find a comma."
)
results = list(stream_chunks(text, generate=gen, clock=clock))
assert any(r.starved for r in results), "expected the gap to be flagged"
assert any(r.drained > 0 for r in results)
def test_full_text_reconstructed():
"""Every sentence is emitted exactly once, in order."""
results = run(true_rtf=3.8)
joined = " ".join(r.text for r in results)
assert joined == " ".join(split_sentences(PARAGRAPH))
def _main():
results = run(true_rtf=3.8)
print(f"{'idx':>3} {'chars':>5} {'audio_s':>8} {'gen_s':>7} "
f"{'est_s':>6} {'buf_before':>10} {'buf_after':>9} {'rtf':>5} {'drain':>6}")
for r in results:
print(f"{r.index:>3} {len(r.text):>5} {r.audio_sec:>8.2f} {r.gen_time:>7.3f} "
f"{r.est_gen:>6.3f} {r.buffer_before:>10.2f} {r.buffer_after:>9.2f} "
f"{r.rtf:>5.2f} {r.drained:>6.3f}")
starved = [r.index for r in results if r.starved]
print(f"\nchunks={len(results)} starved={starved or 'none'} "
f"ttfa≈{results[0].gen_time:.3f}s")
if __name__ == "__main__":
_main()