@@ -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