Skip to content

Commit c564930

Browse files
committed
Add phoneme prediction on exportable .json and .csv files
Signed-off-by: Edresson Casanova <edresson1@gmail.com>
1 parent 1f41871 commit c564930

3 files changed

Lines changed: 75 additions & 2 deletions

File tree

examples/tts/magpietts_inference.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -190,6 +190,8 @@ def _enrich_filewise_metrics_with_manifest(filewise_metrics: list, manifest_path
190190
"rank_local_idx",
191191
"audio_filepath",
192192
"context_audio_filepath",
193+
"predicted_phoneme_text",
194+
"predicted_phoneme_tokens",
193195
]:
194196
if key in record and key not in new_row:
195197
new_row[key] = record[key]
@@ -249,6 +251,8 @@ def turn_sort_key(r):
249251
eou_type_turns = [r.get("eou_type") for r in turns]
250252
eou_trailing_duration_turns = [r.get("eou_trailing_duration") for r in turns]
251253
eou_trail_rms_ratio_turns = [r.get("eou_trail_rms_ratio") for r in turns]
254+
predicted_phoneme_text_turns = [r.get("predicted_phoneme_text", "") for r in turns]
255+
predicted_phoneme_tokens_turns = [r.get("predicted_phoneme_tokens", []) for r in turns]
252256

253257
grouped_rows.append(
254258
{
@@ -280,6 +284,8 @@ def turn_sort_key(r):
280284
"eou_type_turns": eou_type_turns,
281285
"eou_trailing_duration_turns": eou_trailing_duration_turns,
282286
"eou_trail_rms_ratio_turns": eou_trail_rms_ratio_turns,
287+
"predicted_phoneme_text_turns": predicted_phoneme_text_turns,
288+
"predicted_phoneme_tokens_turns": predicted_phoneme_tokens_turns,
283289
"reference_text": [r.get("gt_text", "") for r in turns],
284290
"asr_hyp": [r.get("pred_text", "") for r in turns],
285291
"pred_audio_paths": [r.get("pred_audio_filepath", "") for r in turns],
@@ -327,6 +333,8 @@ def _write_grouped_multiturn_filewise_metrics_csv(csv_path: str, grouped_rows: l
327333
"eou_type_turns",
328334
"eou_trailing_duration_turns",
329335
"eou_trail_rms_ratio_turns",
336+
"predicted_phoneme_text_turns",
337+
"predicted_phoneme_tokens_turns",
330338
"target_audio_path",
331339
"context_audio_path",
332340
"pred_audio_paths",

nemo/collections/tts/modules/magpietts_inference/evaluate_generated_audio.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -89,6 +89,8 @@ def strip_text_annotations_from_text(text: str) -> str:
8989
'pred_context_ssim',
9090
'pred_gt_esim',
9191
'pred_gt_ems',
92+
'predicted_phoneme_text',
93+
'predicted_phoneme_tokens',
9294
'pred_text',
9395
'gt_text',
9496
'gt_audio_filepath',
@@ -690,6 +692,9 @@ def evaluate_dir(
690692
'total_gen_audio_seconds': file_duration,
691693
'predicted_codes_path': codes_file_lists[ridx] if has_codes else None,
692694
}
695+
for manifest_debug_key in ['predicted_phoneme_text', 'predicted_phoneme_tokens']:
696+
if manifest_debug_key in record:
697+
metric_row[manifest_debug_key] = record[manifest_debug_key]
693698
if with_emotion_metrics:
694699
metric_row['pred_gt_esim'] = pred_gt_esim
695700
metric_row['pred_gt_ems'] = pred_gt_ems

nemo/collections/tts/modules/magpietts_inference/inference.py

Lines changed: 62 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1167,6 +1167,40 @@ def _ensure_codec_silence_codes(self) -> torch.Tensor:
11671167

11681168
return self.model._codec_sil_codes_buffer.to(self.model.device).long()
11691169

1170+
def _decode_phoneme_prediction_slice(self, state, start_step: int, end_step: int) -> Tuple[List[int], str]:
1171+
"""Decode predicted phoneme-channel tokens for one generated turn."""
1172+
phoneme_tokenizer = getattr(self.model, "phoneme_tokenizer", None)
1173+
phoneme_predictions = getattr(state, "all_phoneme_predictions", None)
1174+
if phoneme_tokenizer is None or not phoneme_predictions or end_step <= start_step:
1175+
return [], ""
1176+
1177+
phoneme_slice = phoneme_predictions[start_step:end_step]
1178+
if len(phoneme_slice) == 0:
1179+
return [], ""
1180+
1181+
phoneme_tensor = torch.stack(phoneme_slice, dim=-1) # (B, S, T)
1182+
tokens = phoneme_tensor[0].detach().cpu().T.reshape(-1).tolist()
1183+
1184+
special_tokens = {
1185+
getattr(phoneme_tokenizer, "bos_token_id", None),
1186+
getattr(phoneme_tokenizer, "eos_token_id", None),
1187+
getattr(phoneme_tokenizer, "pad_token_id", None),
1188+
getattr(phoneme_tokenizer, "pad", None),
1189+
}
1190+
special_tokens = {int(token) for token in special_tokens if token is not None}
1191+
tokens = [int(token) for token in tokens if int(token) >= 0 and int(token) not in special_tokens]
1192+
1193+
if not tokens:
1194+
return [], ""
1195+
1196+
try:
1197+
phoneme_text = phoneme_tokenizer.decode(tokens)
1198+
except Exception as exc:
1199+
logging.warning(f"Could not decode predicted phoneme tokens for multiturn debug export: {exc}")
1200+
phoneme_text = ""
1201+
1202+
return tokens, phoneme_text
1203+
11701204
def _run_multiturn_generation(self, batch: Dict[str, Any]):
11711205
model = self.model
11721206
device = model.device
@@ -1237,6 +1271,7 @@ def _run_multiturn_generation(self, batch: Dict[str, Any]):
12371271
)
12381272

12391273
turn_frame_ranges = []
1274+
turn_phoneme_outputs = []
12401275
decode_start_frame = 0
12411276
max_decoder_steps = params.max_decoder_steps
12421277

@@ -1357,6 +1392,8 @@ def _run_multiturn_generation(self, batch: Dict[str, Any]):
13571392
user_audio_channel_embedding=user_audio_channel_embedding,
13581393
)
13591394

1395+
turn_phoneme_start_step = len(getattr(state, "all_phoneme_predictions", []))
1396+
13601397
for i in range(delay_tokens):
13611398
user_step_emb = warmup_user_audio[:, i] if warmup_user_audio is not None else None
13621399
state.finished.zero_()
@@ -1405,7 +1442,20 @@ def _run_multiturn_generation(self, batch: Dict[str, Any]):
14051442
state.audio_prediction_end_idx.fill_(-1)
14061443
state.finished.zero_()
14071444
turn_end_frame = sum(p.size(-1) for p in state.all_predictions)
1445+
turn_phoneme_end_step = len(getattr(state, "all_phoneme_predictions", []))
1446+
predicted_phoneme_tokens, predicted_phoneme_text = self._decode_phoneme_prediction_slice(
1447+
state,
1448+
turn_phoneme_start_step,
1449+
turn_phoneme_end_step,
1450+
)
14081451
turn_frame_ranges.append((turn_id, turn_start_frame, turn_end_frame))
1452+
turn_phoneme_outputs.append(
1453+
{
1454+
"turn_id": int(turn_id),
1455+
"predicted_phoneme_text": predicted_phoneme_text,
1456+
"predicted_phoneme_tokens": predicted_phoneme_tokens,
1457+
}
1458+
)
14091459

14101460
codec_sil_codes = self._ensure_codec_silence_codes()
14111461
bos_id = getattr(model, "audio_bos_id", -1)
@@ -1428,7 +1478,7 @@ def _run_multiturn_generation(self, batch: Dict[str, Any]):
14281478

14291479
finalize_output = model.streaming_finalize(state, use_inference_mode=True)
14301480

1431-
return finalize_output, turn_frame_ranges, decode_start_frame, generated_codes
1481+
return finalize_output, turn_frame_ranges, decode_start_frame, generated_codes, turn_phoneme_outputs
14321482

14331483
@staticmethod
14341484
def _save_code_slice(
@@ -1637,9 +1687,15 @@ def _run_multiturn_user_audio_inference(
16371687
continue
16381688

16391689
start_time = time.time()
1640-
output, turn_frame_ranges, decode_start_frame, generated_codes = self._run_multiturn_generation(batch)
1690+
output, turn_frame_ranges, decode_start_frame, generated_codes, turn_phoneme_outputs = (
1691+
self._run_multiturn_generation(batch)
1692+
)
16411693
elapsed = time.time() - start_time
16421694

1695+
turn_phoneme_outputs_by_turn_id = {
1696+
int(item.get("turn_id", -1)): item for item in turn_phoneme_outputs
1697+
}
1698+
16431699
predicted_audio = output.audio.float().detach().cpu()
16441700
predicted_audio_lens = output.audio_len.int().detach().cpu()
16451701
full_len = int(predicted_audio_lens[0].item())
@@ -1712,6 +1768,8 @@ def _run_multiturn_user_audio_inference(
17121768
description=f"target audio fallback/context for sample_idx={sample_idx}, turn_id={turn_id}",
17131769
)
17141770

1771+
phoneme_debug = turn_phoneme_outputs_by_turn_id.get(int(turn_id), {})
1772+
17151773
turn_manifest_records.append(
17161774
{
17171775
"audio_filepath": f"target_audio_{item_idx}.wav",
@@ -1720,6 +1778,8 @@ def _run_multiturn_user_audio_inference(
17201778
"speaker": str(sample_idx),
17211779
"source_sample_idx": sample_idx,
17221780
"turn_id": int(turn_id),
1781+
"predicted_phoneme_text": phoneme_debug.get("predicted_phoneme_text", ""),
1782+
"predicted_phoneme_tokens": phoneme_debug.get("predicted_phoneme_tokens", []),
17231783
}
17241784
)
17251785
logging.info(

0 commit comments

Comments
 (0)