From 8675a5eb258fea6b07d633642bfd7a5a625f6bfa Mon Sep 17 00:00:00 2001 From: pengzhendong <275331498@qq.com> Date: Fri, 25 Apr 2025 03:58:16 +0000 Subject: [PATCH] return wav_path as utt id --- dataset.py | 1 + recognize.py | 4 ++-- speech_llm.py | 5 ++++- 3 files changed, 7 insertions(+), 3 deletions(-) diff --git a/dataset.py b/dataset.py index 02e8c11..92600a4 100644 --- a/dataset.py +++ b/dataset.py @@ -118,6 +118,7 @@ def __getitem__(self, i) -> Dict[str, torch.Tensor]: ctc_ids = ctc_tokens['input_ids'][0] ctc_ids_len = ctc_tokens['attention_mask'].sum().item() ret = { + 'wav': msg['wav'], 'input_ids': input_ids, 'attention_mask': attention_mask, 'mel': mel, diff --git a/recognize.py b/recognize.py index ecdb5f3..c951c3a 100644 --- a/recognize.py +++ b/recognize.py @@ -60,9 +60,9 @@ def main(): text = tokenizer.batch_decode(generated_ids, skip_special_tokens=True) print(text) - for t in text: + for wav, t in zip(item['wav'], text): t = t.replace('\n', ' ') - fid.write(t + '\n') + fid.write(wav + ' ' + t + '\n') sys.stdout.flush() fid.close() diff --git a/speech_llm.py b/speech_llm.py index 17cb4f9..5755ade 100644 --- a/speech_llm.py +++ b/speech_llm.py @@ -1,7 +1,7 @@ # Copyright (c) 2025 Binbin Zhang(binbzha@qq.com) import math -from typing import Optional +from typing import List, Optional from dataclasses import dataclass, field import safetensors @@ -148,6 +148,7 @@ def get_speech_embeddings(self, mel, mel_len): @torch.autocast(device_type="cuda", dtype=torch.bfloat16) def forward( self, + wav: List[str] = None, input_ids: torch.LongTensor = None, attention_mask: Optional[torch.Tensor] = None, labels: Optional[torch.LongTensor] = None, @@ -181,6 +182,7 @@ def forward( @torch.autocast(device_type="cuda", dtype=torch.bfloat16) def generate( self, + wav: List[str] = None, input_ids: torch.LongTensor = None, attention_mask: Optional[torch.Tensor] = None, mel: torch.LongTensor = None, @@ -207,6 +209,7 @@ def generate( @torch.autocast(device_type="cuda", dtype=torch.bfloat16) def decode_ctc( self, + wav: List[str] = None, input_ids: torch.LongTensor = None, attention_mask: Optional[torch.Tensor] = None, mel: torch.LongTensor = None,