shixian.shi
2024-01-11 70379713923e9938dcfcb3791b17f7b469233432
funasr/bin/inference.py
@@ -1,26 +1,26 @@
import os.path
import torch
import numpy as np
import hydra
import json
from omegaconf import DictConfig, OmegaConf, ListConfig
import logging
from funasr.download.download_from_hub import download_model
from funasr.train_utils.set_all_random_seed import set_all_random_seed
from funasr.utils.load_utils import load_bytes
from funasr.train_utils.device_funcs import to_device
from tqdm import tqdm
from funasr.train_utils.load_pretrained_model import load_pretrained_model
import time
import torch
import hydra
import random
import string
from funasr.register import tables
import logging
import os.path
from tqdm import tqdm
from omegaconf import DictConfig, OmegaConf, ListConfig
from funasr.utils.load_utils import load_audio_text_image_video, extract_fbank
from funasr.utils.vad_utils import slice_padding_audio_samples
from funasr.utils.timestamp_tools import time_stamp_sentence
from funasr.register import tables
from funasr.utils.load_utils import load_bytes
from funasr.download.file import download_from_url
from funasr.download.download_from_hub import download_model
from funasr.utils.vad_utils import slice_padding_audio_samples
from funasr.train_utils.set_all_random_seed import set_all_random_seed
from funasr.train_utils.load_pretrained_model import load_pretrained_model
from funasr.utils.load_utils import load_audio_text_image_video, extract_fbank
from funasr.utils.timestamp_tools import timestamp_sentence
from funasr.models.campplus.utils import sv_chunk, postprocess, distribute_spk
from funasr.models.campplus.cluster_backend import ClusterBackend
def prepare_data_iterator(data_in, input_len=None, data_type=None, key=None):
   """
@@ -126,13 +126,27 @@
         punc_kwargs = {"model": punc_model, "model_revision": punc_kwargs}
         punc_model, punc_kwargs = self.build_model(**punc_kwargs)
         
        # if spk_model is not None, build spk model else None
        spk_model = kwargs.get("spk_model", None)
        spk_kwargs = kwargs.get("spk_model_revision", None)
        if spk_model is not None:
            spk_kwargs = {"model": spk_model, "model_revision": spk_kwargs}
            spk_model, spk_kwargs = self.build_model(**spk_kwargs)
            self.cb_model = ClusterBackend()
            spk_mode = kwargs.get("spk_mode", 'punc_segment')
            if spk_mode not in ["default", "vad_segment", "punc_segment"]:
                logging.error("spk_mode should be one of default, vad_segment and punc_segment.")
            self.spk_mode = spk_mode
            logging.warning("Many to print when using speaker model...")
      self.kwargs = kwargs
      self.model = model
      self.vad_model = vad_model
      self.vad_kwargs = vad_kwargs
      self.punc_model = punc_model
      self.punc_kwargs = punc_kwargs
        self.spk_model = spk_model
        self.spk_kwargs = spk_kwargs
      
   def build_model(self, **kwargs):
@@ -198,7 +212,6 @@
         return self.generate_with_vad(input, input_len=input_len, **cfg)
      
   def generate(self, input, input_len=None, model=None, kwargs=None, key=None, **cfg):
      # import pdb; pdb.set_trace()
      kwargs = self.kwargs if kwargs is None else kwargs
      kwargs.update(cfg)
      model = self.model if model is None else model
@@ -260,6 +273,7 @@
      kwargs.update(cfg)
      beg_vad = time.time()
      res = self.generate(input, input_len=input_len, model=model, kwargs=kwargs, **cfg)
        vad_res = res
      end_vad = time.time()
      print(f"time cost vad: {end_vad - beg_vad:0.3f}")
@@ -314,10 +328,20 @@
            batch_size_ms_cum = 0
            end_idx = j + 1
            speech_j, speech_lengths_j = slice_padding_audio_samples(speech, speech_lengths, sorted_data[beg_idx:end_idx])
            beg_idx = end_idx
            results = self.generate(speech_j, input_len=None, model=model, kwargs=kwargs, **cfg)
                if self.spk_model is not None:
                    all_segments = []
                    # compose vad segments: [[start_time_sec, end_time_sec, speech], [...]]
                    for _b in range(len(speech_j)):
                        vad_segments = [[sorted_data[beg_idx:end_idx][_b][0][0]/1000.0, \
                                        sorted_data[beg_idx:end_idx][_b][0][1]/1000.0, \
                                        speech_j[_b]]]
                        segments = sv_chunk(vad_segments)
                        all_segments.extend(segments)
                        speech_b = [i[2] for i in segments]
                        spk_res = self.generate(speech_b, input_len=None, model=self.spk_model, kwargs=kwargs, **cfg)
                        results[_b]['spk_embedding'] = spk_res[0]['spk_embedding']
                beg_idx = end_idx
            if len(results) < 1:
               continue
            results_sorted.extend(results)
@@ -336,39 +360,63 @@
            restored_data[index] = results_sorted[j]
         result = {}
         
            # results combine for texts, timestamps, speaker embeddings and others
            # TODO: rewrite for clean code
         for j in range(n):
            for k, v in restored_data[j].items():
               if not k.startswith("timestamp"):
                    if k.startswith("timestamp"):
                  if k not in result:
                     result[k] = restored_data[j][k]
                  else:
                     result[k] += restored_data[j][k]
               else:
                  result[k] = []
                  for t in restored_data[j][k]:
                     t[0] += vadsegments[j][0]
                     t[1] += vadsegments[j][0]
                        result[k].extend(restored_data[j][k])
                    elif k == 'spk_embedding':
                        if k not in result:
                            result[k] = restored_data[j][k]
                        else:
                            result[k] = torch.cat([result[k], restored_data[j][k]], dim=0)
                    elif k == 'text':
                        if k not in result:
                            result[k] = restored_data[j][k]
                        else:
                            result[k] += " " + restored_data[j][k]
                    else:
                        if k not in result:
                            result[k] = restored_data[j][k]
                        else:
                  result[k] += restored_data[j][k]
            # step.3 compute punc model
            if self.punc_model is not None:
                self.punc_kwargs.update(cfg)
                punc_res = self.generate(result["text"], model=self.punc_model, kwargs=self.punc_kwargs, **cfg)
                result["text_with_punc"] = punc_res[0]["text"]
            # speaker embedding cluster after resorted
            if self.spk_model is not None:
                all_segments = sorted(all_segments, key=lambda x: x[0])
                spk_embedding = result['spk_embedding']
                labels = self.cb_model(spk_embedding)
                del result['spk_embedding']
                sv_output = postprocess(all_segments, None, labels, spk_embedding)
                if self.spk_mode == 'vad_segment':
                    sentence_list = []
                    for res, vadsegment in zip(restored_data, vadsegments):
                        sentence_list.append({"start": vadsegment[0],\
                                                "end": vadsegment[1],
                                                "sentence": res['text'],
                                                "timestamp": res['timestamp']})
                else: # punc_segment
                    sentence_list = timestamp_sentence(punc_res[0]['punc_array'], \
                                                        result['timestamp'], \
                                                        result['text'])
                distribute_spk(sentence_list, sv_output)
                result['sentence_info'] = sentence_list
                  
         result["key"] = key
         results_ret_list.append(result)
         pbar_total.update(1)
      # step.3 compute punc model
      model = self.punc_model
      kwargs = self.punc_kwargs
      kwargs.update(cfg)
      for i, result in enumerate(results_ret_list):
         beg_punc = time.time()
         res = self.generate(result["text"], model=model, kwargs=kwargs, **cfg)
         end_punc = time.time()
         print(f"time punc: {end_punc - beg_punc:0.3f}")
         # sentences = time_stamp_sentence(model.punc_list, model.sentence_end_id, results_ret_list[i]["timestamp"], res[i]["text"])
         # results_ret_list[i]["time_stamp"] = res[0]["text_postprocessed_punc"]
         # results_ret_list[i]["sentences"] = sentences
         results_ret_list[i]["text_with_punc"] = res[i]["text"]
      
      pbar_total.update(1)
      end_total = time.time()