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
+70 -7
View File
@@ -164,9 +164,8 @@ class Engine:
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."""
def _raw_wav(self, text: str, knobs: "TTSRequest") -> torch.Tensor:
"""One generate pass → wav tensor [1,T], CUDA-synced for honest timing."""
with torch.inference_mode():
wav = self.model.generate(
text,
@@ -177,8 +176,30 @@ class Engine:
)
if DEVICE.startswith("cuda"):
torch.cuda.synchronize()
audio_sec = wav.shape[-1] / self.sr
return wav, audio_sec
return wav
def generate(self, text: str, knobs: "TTSRequest") -> tuple[torch.Tensor, float]:
"""Synthesize ``text`` → (wav [1,T], audio_seconds)."""
wav = self._raw_wav(text, knobs)
return wav, wav.shape[-1] / self.sr
def generate_primed(
self, content: str, context: str, knobs: "TTSRequest"
) -> tuple[torch.Tensor, float]:
"""Context-primed generate (plan §1.6): render ``context + content``
together so ``content``'s prosody knows what preceded it, then discard the
context audio. The cut snaps to the inter-sentence pause (energy minimum)
near the context's solo duration, with a short fade-in to kill any click.
Costs two passes (context-solo to locate the cut, then the joint) — the
scheduler budgets for this via prime_cost_factor."""
ctx_wav = self._raw_wav(context, knobs)
d_ctx = ctx_wav.shape[-1] / self.sr
joint = self._raw_wav(f"{context} {content}", knobs)
cut = _cut_at_pause(joint, self.sr, d_ctx)
kept = joint[..., cut:].clone()
_fade_in(kept, self.sr)
return kept, kept.shape[-1] / self.sr
engine = Engine()
@@ -206,12 +227,50 @@ class TTSRequest(BaseModel):
top_p: float = 0.95
top_k: int = 1000
repetition_penalty: float = 1.2
# Context-priming at joins (plan §1.6) — opt-in for A/B.
prime: bool = False
prime_first_n: int = 2 # how many early joins to prime when prime=True
# 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 _cut_at_pause(
joint: torch.Tensor, sr: int, approx_sec: float,
window_sec: float = 0.4, frame_sec: float = 0.02,
) -> int:
"""Sample index to cut the discarded context off ``joint``. Searches ±window
around the context's solo duration for the lowest-energy 20 ms frame — the
inter-sentence pause — so the seam lands in silence, not mid-phone."""
audio = joint.reshape(-1)
n = audio.shape[0]
center = int(approx_sec * sr)
w = int(window_sec * sr)
frame = max(1, int(frame_sec * sr))
lo = max(0, center - w)
hi = min(n - frame, center + w)
if hi <= lo:
return min(center, n)
best_i, best_e = center, float("inf")
i = lo
while i < hi:
e = float((audio[i:i + frame] ** 2).mean())
if e < best_e:
best_e, best_i = e, i
i += frame
return best_i
def _fade_in(wav: torch.Tensor, sr: int, fade_sec: float = 0.005) -> None:
"""In-place short fade-in to remove any click at the cut seam."""
fade = int(fade_sec * sr)
if wav.shape[-1] > fade > 0:
ramp = torch.linspace(0.0, 1.0, fade, device=wav.device, dtype=wav.dtype)
wav[..., :fade] *= ramp
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()
@@ -234,6 +293,8 @@ def _chunk_config(req: TTSRequest) -> ChunkConfig:
cfg.margin_first = req.margin_first
if req.rtf_prior is not None:
cfg.rtf_prior = req.rtf_prior
if req.prime:
cfg.prime_first_n = req.prime_first_n
return cfg
@@ -289,7 +350,9 @@ def tts(req: TTSRequest) -> StreamingResponse:
yield _pcm16(wav)
return
def _gen(text: str) -> tuple[torch.Tensor, float]:
def _gen(text: str, context: str | None) -> tuple[torch.Tensor, float]:
if context:
return engine.generate_primed(text, context, req)
return engine.generate(text, req)
total_audio = 0.0
@@ -305,7 +368,7 @@ def tts(req: TTSRequest) -> StreamingResponse:
def _log_chunk(r: ChunkResult, ttfa_ms: float) -> None:
tag = " STARVED" if r.starved else ""
tag = (" PRIMED" if r.primed else "") + (" 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)