voidful commited on
Commit
5f2a1b5
·
1 Parent(s): 340bbec

Add active-voiced pace correction ablation

Browse files
Files changed (4) hide show
  1. app.py +22 -1
  2. production.py +47 -0
  3. tests/test_production.py +73 -0
  4. tests/test_release_pins.py +35 -0
app.py CHANGED
@@ -18,6 +18,7 @@ from transformers import PreTrainedTokenizerFast
18
  from bluemagpie import BlueMagpieModel
19
  from production import (
20
  StopHysteresisController,
 
21
  apply_loudness_floor,
22
  count_speech_units,
23
  effective_generation_cfg,
@@ -323,7 +324,27 @@ def _generate_chunk(
323
  target_cps=TARGET_CPS,
324
  min_speed=MIN_PACE_SPEED,
325
  )
326
- return _apply_speed(audio, pace_speed)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
327
 
328
 
329
  def _speaker_anchor_array(centroid: torch.Tensor) -> np.ndarray:
 
18
  from bluemagpie import BlueMagpieModel
19
  from production import (
20
  StopHysteresisController,
21
+ active_pace_correction_speed,
22
  apply_loudness_floor,
23
  count_speech_units,
24
  effective_generation_cfg,
 
324
  target_cps=TARGET_CPS,
325
  min_speed=MIN_PACE_SPEED,
326
  )
327
+ waveform = _apply_speed(audio, pace_speed)
328
+ try:
329
+ active_duration = active_voiced_duration_seconds(waveform, SR)
330
+ except (TypeError, ValueError, RuntimeError, OverflowError):
331
+ # Invalid pace evidence must never turn into an unbounded correction.
332
+ # The unchanged waveform will still face the ordinary local gate.
333
+ return waveform
334
+ active_speed = active_pace_correction_speed(
335
+ active_duration,
336
+ text,
337
+ target_cps=TARGET_CPS,
338
+ prior_speed=pace_speed,
339
+ min_total_speed=MIN_PACE_SPEED,
340
+ )
341
+ if active_speed < 1.0:
342
+ print(
343
+ "[BlueMagpie] active pace correction "
344
+ f"active_duration={active_duration:.3f} rate={active_speed:.6f} "
345
+ f"combined_rate={pace_speed * active_speed:.6f}"
346
+ )
347
+ return _apply_speed(waveform, active_speed)
348
 
349
 
350
  def _speaker_anchor_array(centroid: torch.Tensor) -> np.ndarray:
production.py CHANGED
@@ -968,6 +968,53 @@ def target_pace_speed(
968
  return min(1.0, max(float(min_speed), actual_seconds / target_seconds))
969
 
970
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
971
  def duration_hard_stop_steps(
972
  expected_steps: int,
973
  *,
 
968
  return min(1.0, max(float(min_speed), actual_seconds / target_seconds))
969
 
970
 
971
+ def active_pace_correction_speed(
972
+ active_duration_seconds: float,
973
+ text: str,
974
+ *,
975
+ target_cps: float,
976
+ prior_speed: float = 1.0,
977
+ min_total_speed: float = 0.80,
978
+ ) -> float:
979
+ """Return a second stretch rate from active-voice duration.
980
+
981
+ ``prior_speed`` is the rate already applied by ``target_pace_speed``.
982
+ Successive pitch-preserving stretch rates multiply, so the second rate is
983
+ floored at ``min_total_speed / prior_speed``. Invalid evidence and audio
984
+ that is already at or below the target active CPS fail safely to ``1.0``;
985
+ this helper never speeds audio up.
986
+ """
987
+
988
+ units = count_speech_units(text)
989
+ try:
990
+ active_seconds = float(active_duration_seconds)
991
+ target = float(target_cps)
992
+ previous_rate = float(prior_speed)
993
+ total_floor = float(min_total_speed)
994
+ except (TypeError, ValueError, OverflowError):
995
+ return 1.0
996
+ if (
997
+ units <= 0
998
+ or not math.isfinite(active_seconds)
999
+ or active_seconds <= 0.0
1000
+ or not math.isfinite(target)
1001
+ or target <= 0.0
1002
+ or not math.isfinite(previous_rate)
1003
+ or not 0.0 < previous_rate <= 1.0
1004
+ or not math.isfinite(total_floor)
1005
+ or not 0.0 < total_floor <= 1.0
1006
+ ):
1007
+ return 1.0
1008
+
1009
+ desired_rate = active_seconds * target / float(units)
1010
+ if not math.isfinite(desired_rate) or desired_rate >= 1.0:
1011
+ return 1.0
1012
+ remaining_floor = total_floor / previous_rate
1013
+ if remaining_floor >= 1.0:
1014
+ return 1.0
1015
+ return min(1.0, max(remaining_floor, desired_rate))
1016
+
1017
+
1018
  def duration_hard_stop_steps(
1019
  expected_steps: int,
1020
  *,
tests/test_production.py CHANGED
@@ -7,6 +7,7 @@ from torch import nn
7
 
8
  from production import (
9
  StopHysteresisController,
 
10
  candidate_local_score,
11
  candidate_transition_score,
12
  compare_asr_text,
@@ -226,6 +227,78 @@ def test_target_pace_speed_slows_completed_audio_without_extending_generation():
226
  assert target_pace_speed(2500, 1000, text, target_cps=4.0) == 1.0
227
 
228
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
229
  def test_finish_audio_fades_endpoint_and_appends_silence():
230
  output = finish_audio(
231
  np.ones(10, dtype=np.float32),
 
7
 
8
  from production import (
9
  StopHysteresisController,
10
+ active_pace_correction_speed,
11
  candidate_local_score,
12
  candidate_transition_score,
13
  compare_asr_text,
 
227
  assert target_pace_speed(2500, 1000, text, target_cps=4.0) == 1.0
228
 
229
 
230
+ def test_active_pace_correction_uses_active_duration_formula():
231
+ text = "一二三四五六七八"
232
+
233
+ assert active_pace_correction_speed(
234
+ 1.8,
235
+ text,
236
+ target_cps=4.0,
237
+ prior_speed=1.0,
238
+ ) == pytest.approx(0.9)
239
+
240
+
241
+ def test_active_pace_correction_never_speeds_up_slow_audio():
242
+ assert active_pace_correction_speed(
243
+ 2.2,
244
+ "一二三四五六七八",
245
+ target_cps=4.0,
246
+ prior_speed=1.0,
247
+ ) == 1.0
248
+
249
+
250
+ def test_active_pace_correction_respects_combined_stretch_cap():
251
+ text = "一二三四五六七八"
252
+ second_rate = active_pace_correction_speed(
253
+ 1.0,
254
+ text,
255
+ target_cps=4.0,
256
+ prior_speed=0.9,
257
+ min_total_speed=0.8,
258
+ )
259
+
260
+ assert second_rate == pytest.approx(0.8 / 0.9)
261
+ assert 0.9 * second_rate == pytest.approx(0.8)
262
+ assert active_pace_correction_speed(
263
+ 1.0,
264
+ text,
265
+ target_cps=4.0,
266
+ prior_speed=0.8,
267
+ min_total_speed=0.8,
268
+ ) == 1.0
269
+
270
+
271
+ @pytest.mark.parametrize(
272
+ ("active_duration", "text", "target_cps", "prior_speed", "minimum"),
273
+ [
274
+ (0.0, "完整文字", 4.0, 1.0, 0.8),
275
+ (-1.0, "完整文字", 4.0, 1.0, 0.8),
276
+ (float("nan"), "完整文字", 4.0, 1.0, 0.8),
277
+ (float("inf"), "完整文字", 4.0, 1.0, 0.8),
278
+ (1.0, "", 4.0, 1.0, 0.8),
279
+ (1.0, "完整文字", 0.0, 1.0, 0.8),
280
+ (1.0, "完整文字", 4.0, 0.0, 0.8),
281
+ (1.0, "完整文字", 4.0, float("nan"), 0.8),
282
+ (1.0, "完整文字", 4.0, 1.0, 0.0),
283
+ (1.0, "完整文字", 4.0, 1.0, float("nan")),
284
+ ],
285
+ )
286
+ def test_active_pace_correction_invalid_evidence_fails_safe(
287
+ active_duration,
288
+ text,
289
+ target_cps,
290
+ prior_speed,
291
+ minimum,
292
+ ):
293
+ assert active_pace_correction_speed(
294
+ active_duration,
295
+ text,
296
+ target_cps=target_cps,
297
+ prior_speed=prior_speed,
298
+ min_total_speed=minimum,
299
+ ) == 1.0
300
+
301
+
302
  def test_finish_audio_fades_endpoint_and_appends_silence():
303
  output = finish_audio(
304
  np.ones(10, dtype=np.float32),
tests/test_release_pins.py CHANGED
@@ -242,6 +242,41 @@ def test_app_reverifies_the_post_join_speed_adjusted_whole_waveform():
242
  assert "RELEASE_SPEAKER_TRIGGER_SECONDS" in source
243
 
244
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
245
  def test_whole_candidate_qualification_uses_the_exact_return_assembler_after_local_pass():
246
  source = (ROOT / "app.py").read_text(encoding="utf-8")
247
  tree = ast.parse(source)
 
242
  assert "RELEASE_SPEAKER_TRIGGER_SECONDS" in source
243
 
244
 
245
+ def test_chunk_generation_corrects_active_pace_after_total_duration_pace():
246
+ source = (ROOT / "app.py").read_text(encoding="utf-8")
247
+ tree = ast.parse(source)
248
+ functions = {
249
+ node.name: node
250
+ for node in tree.body
251
+ if isinstance(node, ast.FunctionDef)
252
+ }
253
+ generate_source = ast.get_source_segment(source, functions["_generate_chunk"])
254
+
255
+ assert generate_source is not None
256
+ total_pace_index = generate_source.index("pace_speed = target_pace_speed(")
257
+ first_stretch_index = generate_source.index(
258
+ "waveform = _apply_speed(audio, pace_speed)"
259
+ )
260
+ active_measure_index = generate_source.index(
261
+ "active_voiced_duration_seconds(waveform, SR)"
262
+ )
263
+ active_pace_index = generate_source.index(
264
+ "active_speed = active_pace_correction_speed("
265
+ )
266
+ second_stretch_index = generate_source.index(
267
+ "return _apply_speed(waveform, active_speed)"
268
+ )
269
+ assert (
270
+ total_pace_index
271
+ < first_stretch_index
272
+ < active_measure_index
273
+ < active_pace_index
274
+ < second_stretch_index
275
+ )
276
+ assert "prior_speed=pace_speed" in generate_source
277
+ assert "min_total_speed=MIN_PACE_SPEED" in generate_source
278
+
279
+
280
  def test_whole_candidate_qualification_uses_the_exact_return_assembler_after_local_pass():
281
  source = (ROOT / "app.py").read_text(encoding="utf-8")
282
  tree = ast.parse(source)