feat(chatterbox-fast): context-priming at joins (§1.6, opt-in)

Prime early joins by prepending the prior sentence as backward prosodic context,
generating context+content together, then discarding the context audio. The cut
snaps to the inter-sentence pause (energy-minimum search around the context's
solo duration) with a 5ms fade-in to kill any seam click (app: _cut_at_pause /
_fade_in / Engine.generate_primed). Opt-in via request `prime` (default off).

Scheduler: priming is AFFORDABILITY-GATED so it can never starve. A primed chunk
costs ~(2·context + content)/rtf (a 2nd context-solo pass); a chunk is only primed
when buffer ≥ prime_buffer_factor (1.5) × that cost, else it falls back to a cold
generate. Consequences proven in the GPU-free sim (17 tests):
  - fires on early joins for any GPU at/above rtf_prior (3.4 = 3090; A6000 ~3.8-4.0)
  - self-skips (degrades to cold) on a slower-than-fleet GPU rather than starving
  - never primes chunk 0 (latency-critical)
Also fixed a latent Phase-1 bug: margin_first was applied at chunk 0 (budget always
0 there) so it never did anything — now applied at chunk 1 (the first transition).

Live A/B on irv-ml1 (A6000, GLaDOS): TTFB unaffected (445 vs 467ms), no starvation;
priming fired on chunk 2 (gen 1.6s for the doubled pass). On typical text exactly
ONE early join safely primes — priming chunk 2 flattens the buffer so later/larger
chunks no longer clear the safety gate. Samples: ~/chatterbox-ab/_p2_{cold,primed}.wav.
This commit is contained in:
vh
2026-06-01 23:11:21 -07:00
parent 3a92fcd943
commit d707439041
5 changed files with 193 additions and 22 deletions
+64 -7
View File
@@ -63,6 +63,24 @@ class ChunkConfig:
# granularity — plan §1.1).
max_first_sec: float = 2.0
# Context-priming at joins (plan §1.6). Prime the first N joins (chunks
# 1..N) by prepending the prior sentence as backward prosodic context, then
# discarding its audio. 0 ⇒ off. Priming runs a 2nd "context-solo" generate,
# so a primed chunk costs ~(2·context + content)/rtf. Priming is AFFORDABILITY-
# GATED: a chunk is only primed when that cost fits the buffer; otherwise it
# falls back to a cold (unprimed) generate, so priming can never starve the
# stream. The earliest joins (smallest buffer) thus self-skip until the buffer
# has ratcheted up enough to pay for the extra pass.
prime_first_n: int = 0
# Priming headroom: only prime when the buffer is at least this multiple of
# the primed cost, so the 2nd pass doesn't flatten the buffer below the slack
# the NEXT chunk needs to absorb RTF-estimate error. At 1.5, priming fires on
# the early joins for any GPU at/above the rtf_prior floor (3.4 = 3090; A6000
# ~3.8–4.0), and on a slower-than-fleet GPU it self-skips entirely (degrades
# to cold/unprimed) rather than starving.
prime_buffer_factor: float = 1.5
# ── result record ─────────────────────────────────────────────────────────
@@ -82,6 +100,7 @@ class ChunkResult:
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
primed: bool = False # context-priming was applied to this chunk
@property
def starved(self) -> bool:
@@ -137,6 +156,11 @@ def _est_gen_time(text: str, *, rtf: float, sec_per_char: float) -> float:
return (len(text) * sec_per_char) / rtf
def _est_primed_gen_time(content: str, context: str, *, rtf: float, sec_per_char: float) -> float:
"""Primed cost = context-solo pass + joint(context+content) pass."""
return ((2 * len(context) + len(content)) * sec_per_char) / rtf
def plan_chunk(
remaining: Sequence[str],
buffer_remaining: float,
@@ -144,19 +168,27 @@ def plan_chunk(
margin: float,
rtf: float,
sec_per_char: float,
prime_context: str | None = None,
) -> 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.
When ``prime_context`` is set the chunk will be context-primed, so packing
uses the (larger) primed cost estimate to leave room for the 2nd pass.
"""
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:
if prime_context is not None:
est = _est_primed_gen_time(candidate, prime_context, rtf=rtf, sec_per_char=sec_per_char)
else:
est = _est_gen_time(candidate, rtf=rtf, sec_per_char=sec_per_char)
if est > budget:
break
chunk.append(remaining[i])
i += 1
@@ -194,8 +226,10 @@ def _ema(old: float, new: float, alpha: float) -> float:
# ── the online loop ───────────────────────────────────────────────────────
# generate(text) -> (audio_payload, audio_seconds)
GenerateFn = Callable[[str], "tuple[object, float]"]
# generate(text, context) -> (audio_payload, audio_seconds)
# context is the prior sentence to prime backward prosody (discarded by the
# generator), or None for an unprimed chunk.
GenerateFn = Callable[[str, "str | None"], "tuple[object, float]"]
ClockFn = Callable[[], float]
@@ -227,6 +261,7 @@ def stream_chunks(
buffer_remaining = 0.0
remaining: list[str] = units
index = 0
prev_text: str | None = None
while remaining:
first = index == 0
@@ -236,14 +271,34 @@ def stream_chunks(
remaining = relieve_leader(
remaining, buffer_remaining, rtf=rtf, sec_per_char=sec_per_char
)
margin = cfg.margin_first if first else cfg.margin
# margin_first tightens the FIRST transition (planning chunk 1 off chunk
# 0's small buffer — highest starvation risk, plan §1.5).
margin = cfg.margin_first if index == 1 else cfg.margin
# Context-priming (plan §1.6): eligible on chunks 1..N, affordability-gated
# so the 2nd pass can never starve the buffer — fall back to cold otherwise.
context: str | None = None
if 0 < index <= cfg.prime_first_n and prev_text:
prev_units = split_sentences(prev_text)
candidate_ctx = prev_units[-1] if prev_units else prev_text
min_primed = _est_primed_gen_time(
remaining[0], candidate_ctx, rtf=rtf, sec_per_char=sec_per_char
)
if min_primed * cfg.prime_buffer_factor <= buffer_remaining:
context = candidate_ctx
chunk_text, remaining = plan_chunk(
remaining, buffer_remaining, margin=margin, rtf=rtf, sec_per_char=sec_per_char
remaining, buffer_remaining, margin=margin, rtf=rtf,
sec_per_char=sec_per_char, prime_context=context,
)
est_gen = _est_gen_time(chunk_text, rtf=rtf, sec_per_char=sec_per_char)
if context is not None:
est_gen = _est_primed_gen_time(chunk_text, context, rtf=rtf, sec_per_char=sec_per_char)
else:
est_gen = _est_gen_time(chunk_text, rtf=rtf, sec_per_char=sec_per_char)
primed = context is not None
t0 = clock()
audio, audio_sec = generate(chunk_text)
audio, audio_sec = generate(chunk_text, context)
gen_time = clock() - t0
# Starvation: did the buffer run dry while we generated this chunk?
@@ -273,5 +328,7 @@ def stream_chunks(
drained=drained,
rtf=rtf,
sec_per_char=sec_per_char,
primed=primed,
)
prev_text = chunk_text
index += 1