revert(chatterbox-fast): drop context-priming (§1.6) — discard-cut leaks context
Revert the priming feature fromd707439. Live A/B caught an audible artifact: the context-priming discard-cut left part of the throwaway prefix in the output, so a clause ("...without a trace of sarcasm,") was spoken an extra time. Root cause is structural: generate() returns one finished waveform with no marker for where the prefix ends, and the model renders the same prefix with different timing when followed by content than when generated solo — so the duration-estimate + energy-minimum cut is a guess and can leave a sliver (or a whole clause) of prefix in. A reliable cut would need token-level access (the abandoned native-streaming arc) or a per-chunk ASR/alignment pass (heavy, still imperfect, eats the latency budget). Fails the agreed bar: "keep only if it closes the gap without a seam." Kept fromd707439: the .gitignore (build artifacts). NOT re-applied: the bundled margin_first fix — wiring it would shrink chunk 1 (more joins = worse coherence), against the operator's priority, and margin=0.8 there is already starvation-safe. Coherence loss at joins stays an accepted limitation; cold streaming was judged "really good". Phase 1 + Phase 2 parity/perf untouched. Next: Phase 3 deploy.
This commit is contained in:
@@ -164,8 +164,9 @@ class Engine:
|
||||
if DEVICE.startswith("cuda"):
|
||||
torch.cuda.synchronize()
|
||||
|
||||
def _raw_wav(self, text: str, knobs: "TTSRequest") -> torch.Tensor:
|
||||
"""One generate pass → wav tensor [1,T], CUDA-synced for honest timing."""
|
||||
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,
|
||||
@@ -176,30 +177,8 @@ class Engine:
|
||||
)
|
||||
if DEVICE.startswith("cuda"):
|
||||
torch.cuda.synchronize()
|
||||
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
|
||||
audio_sec = wav.shape[-1] / self.sr
|
||||
return wav, audio_sec
|
||||
|
||||
|
||||
engine = Engine()
|
||||
@@ -227,50 +206,12 @@ 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()
|
||||
@@ -293,8 +234,6 @@ 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
|
||||
|
||||
|
||||
@@ -350,9 +289,7 @@ def tts(req: TTSRequest) -> StreamingResponse:
|
||||
yield _pcm16(wav)
|
||||
return
|
||||
|
||||
def _gen(text: str, context: str | None) -> tuple[torch.Tensor, float]:
|
||||
if context:
|
||||
return engine.generate_primed(text, context, req)
|
||||
def _gen(text: str) -> tuple[torch.Tensor, float]:
|
||||
return engine.generate(text, req)
|
||||
|
||||
total_audio = 0.0
|
||||
@@ -368,7 +305,7 @@ def tts(req: TTSRequest) -> StreamingResponse:
|
||||
|
||||
|
||||
def _log_chunk(r: ChunkResult, ttfa_ms: float) -> None:
|
||||
tag = (" PRIMED" if r.primed else "") + (" STARVED" if r.starved else "")
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user