Initial commit: FunASR Speech Recognition Toolkit
Update API Documentation / build-api-docs (push) Has been cancelled
Update API Documentation / build-api-docs (push) Has been cancelled
Add complete FunASR codebase including models, runtime, and documentation.
This commit is contained in:
@@ -0,0 +1,330 @@
|
||||
import torch
|
||||
import random
|
||||
|
||||
|
||||
from funasr.register import tables
|
||||
from funasr.utils.load_utils import extract_fbank, load_audio_text_image_video
|
||||
|
||||
|
||||
@tables.register("dataset_classes", "AudioDataset")
|
||||
class AudioDataset(torch.utils.data.Dataset):
|
||||
"""
|
||||
AudioDataset
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path,
|
||||
index_ds: str = None,
|
||||
frontend=None,
|
||||
tokenizer=None,
|
||||
is_training: bool = True,
|
||||
int_pad_value: int = -1,
|
||||
float_pad_value: float = 0.0,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize AudioDataset.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
index_ds: TODO.
|
||||
frontend: Audio frontend for feature extraction.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
is_training: Boolean flag for training.
|
||||
int_pad_value: TODO.
|
||||
float_pad_value: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
index_ds_class = tables.index_ds_classes.get(index_ds)
|
||||
self.index_ds = index_ds_class(path, **kwargs)
|
||||
|
||||
self.preprocessor_speech = None
|
||||
self.preprocessor_text = None
|
||||
|
||||
if is_training:
|
||||
preprocessor_speech = kwargs.get("preprocessor_speech", None)
|
||||
if preprocessor_speech:
|
||||
preprocessor_speech_class = tables.preprocessor_classes.get(preprocessor_speech)
|
||||
preprocessor_speech = preprocessor_speech_class(
|
||||
**kwargs.get("preprocessor_speech_conf")
|
||||
)
|
||||
self.preprocessor_speech = preprocessor_speech
|
||||
preprocessor_text = kwargs.get("preprocessor_text", None)
|
||||
if preprocessor_text:
|
||||
preprocessor_text_class = tables.preprocessor_classes.get(preprocessor_text)
|
||||
preprocessor_text = preprocessor_text_class(**kwargs.get("preprocessor_text_conf"))
|
||||
self.preprocessor_text = preprocessor_text
|
||||
|
||||
self.frontend = frontend
|
||||
self.fs = 16000 if frontend is None else frontend.fs
|
||||
self.data_type = "sound"
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.int_pad_value = int_pad_value
|
||||
self.float_pad_value = float_pad_value
|
||||
|
||||
def get_source_len(self, index):
|
||||
"""Get source len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_source_len(item)
|
||||
|
||||
def get_target_len(self, index):
|
||||
"""Get target len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_target_len(item)
|
||||
|
||||
def __len__(self):
|
||||
"""Internal: len ."""
|
||||
return len(self.index_ds)
|
||||
|
||||
def __getitem__(self, index):
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
# import pdb;
|
||||
# pdb.set_trace()
|
||||
source = item["source"]
|
||||
data_src = load_audio_text_image_video(source, fs=self.fs)
|
||||
if self.preprocessor_speech:
|
||||
data_src = self.preprocessor_speech(data_src, fs=self.fs)
|
||||
|
||||
speech, speech_lengths = extract_fbank(
|
||||
data_src, data_type=self.data_type, frontend=self.frontend, is_final=True
|
||||
) # speech: [b, T, d]
|
||||
|
||||
target = item["target"]
|
||||
if self.preprocessor_text:
|
||||
target = self.preprocessor_text(target)
|
||||
|
||||
if self.tokenizer:
|
||||
ids = self.tokenizer.encode(target)
|
||||
text = torch.tensor(ids, dtype=torch.int64)
|
||||
else:
|
||||
ids = target
|
||||
text = ids
|
||||
ids_lengths = len(ids)
|
||||
text_lengths = torch.tensor([ids_lengths], dtype=torch.int32)
|
||||
|
||||
return {
|
||||
"speech": speech[0, :, :],
|
||||
"speech_lengths": speech_lengths,
|
||||
"text": text,
|
||||
"text_lengths": text_lengths,
|
||||
}
|
||||
|
||||
def collator(self, samples: list = None):
|
||||
"""Collator.
|
||||
|
||||
Args:
|
||||
samples: TODO.
|
||||
"""
|
||||
outputs = {}
|
||||
for sample in samples:
|
||||
for key in sample.keys():
|
||||
if key not in outputs:
|
||||
outputs[key] = []
|
||||
outputs[key].append(sample[key])
|
||||
|
||||
for key, data_list in outputs.items():
|
||||
if isinstance(data_list[0], torch.Tensor):
|
||||
if data_list[0].dtype == torch.int64 or data_list[0].dtype == torch.int32:
|
||||
|
||||
pad_value = self.int_pad_value
|
||||
else:
|
||||
pad_value = self.float_pad_value
|
||||
|
||||
outputs[key] = torch.nn.utils.rnn.pad_sequence(
|
||||
data_list, batch_first=True, padding_value=pad_value
|
||||
)
|
||||
return outputs
|
||||
|
||||
|
||||
@tables.register("dataset_classes", "AudioDatasetHotword")
|
||||
class AudioDatasetHotword(AudioDataset):
|
||||
# for finetuning contextual_paraformer and seaco_paraformer
|
||||
def __init__(
|
||||
self,
|
||||
*args,
|
||||
seaco_id: bool = 0,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize AudioDatasetHotword.
|
||||
|
||||
Args:
|
||||
*args: Variable positional arguments.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__(*args, **kwargs)
|
||||
self.seaco_id = seaco_id
|
||||
|
||||
def __getitem__(self, index):
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
# import pdb;
|
||||
# pdb.set_trace()
|
||||
source = item["source"]
|
||||
data_src = load_audio_text_image_video(source, fs=self.fs)
|
||||
if self.preprocessor_speech:
|
||||
data_src = self.preprocessor_speech(data_src, fs=self.fs)
|
||||
speech, speech_lengths = extract_fbank(
|
||||
data_src, data_type=self.data_type, frontend=self.frontend, is_final=True
|
||||
) # speech: [b, T, d]
|
||||
|
||||
target = item["target"]
|
||||
if self.preprocessor_text:
|
||||
target = self.preprocessor_text(target)
|
||||
if self.tokenizer:
|
||||
ids = self.tokenizer.encode(target)
|
||||
text = torch.tensor(ids, dtype=torch.int64)
|
||||
else:
|
||||
ids = target
|
||||
text = ids
|
||||
ids_lengths = len(ids)
|
||||
text_lengths = torch.tensor([ids_lengths], dtype=torch.int32)
|
||||
|
||||
def generate_index(
|
||||
length,
|
||||
hotword_min_length=2,
|
||||
hotword_max_length=8,
|
||||
sample_rate=0.75,
|
||||
double_rate=0.1,
|
||||
pre_prob=0.0,
|
||||
pre_index=None,
|
||||
pre_hwlist=None,
|
||||
):
|
||||
"""Generate index.
|
||||
|
||||
Args:
|
||||
length: TODO.
|
||||
hotword_min_length: TODO.
|
||||
hotword_max_length: TODO.
|
||||
sample_rate: TODO.
|
||||
double_rate: TODO.
|
||||
pre_prob: TODO.
|
||||
pre_index: TODO.
|
||||
pre_hwlist: TODO.
|
||||
"""
|
||||
if length < hotword_min_length:
|
||||
return [-1]
|
||||
if random.random() < sample_rate:
|
||||
if pre_prob > 0 and random.random() < pre_prob and pre_index is not None:
|
||||
return pre_index
|
||||
if length == hotword_min_length:
|
||||
return [0, length - 1]
|
||||
elif (
|
||||
random.random() < double_rate
|
||||
and length > hotword_max_length + hotword_min_length + 2
|
||||
):
|
||||
# sample two hotwords in a sentence
|
||||
_max_hw_length = min(hotword_max_length, length // 2)
|
||||
# first hotword
|
||||
start1 = random.randint(0, length // 3)
|
||||
end1 = random.randint(
|
||||
start1 + hotword_min_length - 1, start1 + _max_hw_length - 1
|
||||
)
|
||||
# second hotword
|
||||
start2 = random.randint(end1 + 1, length - hotword_min_length)
|
||||
end2 = random.randint(
|
||||
min(length - 1, start2 + hotword_min_length - 1),
|
||||
min(length - 1, start2 + hotword_max_length - 1),
|
||||
)
|
||||
return [start1, end1, start2, end2]
|
||||
else: # single hotword
|
||||
start = random.randint(0, length - hotword_min_length)
|
||||
end = random.randint(
|
||||
min(length - 1, start + hotword_min_length - 1),
|
||||
min(length - 1, start + hotword_max_length - 1),
|
||||
)
|
||||
return [start, end]
|
||||
else:
|
||||
return [-1]
|
||||
|
||||
hotword_indx = generate_index(text_lengths[0])
|
||||
return {
|
||||
"speech": speech[0, :, :],
|
||||
"speech_lengths": speech_lengths,
|
||||
"text": text,
|
||||
"text_lengths": text_lengths,
|
||||
"hotword_indx": hotword_indx,
|
||||
"seaco_id": self.seaco_id,
|
||||
}
|
||||
|
||||
def collator(self, samples: list = None):
|
||||
"""Collator.
|
||||
|
||||
Args:
|
||||
samples: TODO.
|
||||
"""
|
||||
outputs = {}
|
||||
hotword_indxs = []
|
||||
seaco_id = samples[0]["seaco_id"]
|
||||
for sample in samples:
|
||||
for key in sample.keys():
|
||||
if key == "seaco_id":
|
||||
continue
|
||||
elif key == "hotword_indx":
|
||||
hotword_indxs.append(sample[key])
|
||||
else:
|
||||
if key not in outputs:
|
||||
outputs[key] = []
|
||||
outputs[key].append(sample[key])
|
||||
|
||||
for key, data_list in outputs.items():
|
||||
if isinstance(data_list[0], torch.Tensor):
|
||||
if data_list[0].dtype == torch.int64 or data_list[0].dtype == torch.int32:
|
||||
pad_value = self.int_pad_value
|
||||
else:
|
||||
pad_value = self.float_pad_value
|
||||
outputs[key] = torch.nn.utils.rnn.pad_sequence(
|
||||
data_list, batch_first=True, padding_value=pad_value
|
||||
)
|
||||
|
||||
hotword_list, hotword_lengths = [], []
|
||||
text = outputs["text"]
|
||||
seaco_label_pad = torch.ones_like(text) * -1 if seaco_id else None
|
||||
for b, (hotword_indx, one_text, length) in enumerate(
|
||||
zip(hotword_indxs, text, outputs["text_lengths"])
|
||||
):
|
||||
length = length[0]
|
||||
if seaco_label_pad is not None:
|
||||
seaco_label_pad[b][:length] = seaco_id
|
||||
if hotword_indx[0] != -1:
|
||||
start, end = int(hotword_indx[0]), int(hotword_indx[1])
|
||||
hotword = one_text[start : end + 1]
|
||||
hotword_list.append(hotword)
|
||||
hotword_lengths.append(end - start + 1)
|
||||
if seaco_label_pad is not None:
|
||||
seaco_label_pad[b][start : end + 1] = one_text[start : end + 1]
|
||||
if len(hotword_indx) == 4 and hotword_indx[2] != -1:
|
||||
# the second hotword if exist
|
||||
start, end = int(hotword_indx[2]), int(hotword_indx[3])
|
||||
hotword_list.append(one_text[start : end + 1])
|
||||
hotword_lengths.append(end - start + 1)
|
||||
if seaco_label_pad is not None:
|
||||
seaco_label_pad[b][start : end + 1] = one_text[start : end + 1]
|
||||
hotword_list.append(torch.tensor([1]))
|
||||
hotword_lengths.append(1)
|
||||
hotword_pad = torch.nn.utils.rnn.pad_sequence(
|
||||
hotword_list, batch_first=True, padding_value=0
|
||||
)
|
||||
outputs["hotword_pad"] = hotword_pad
|
||||
outputs["hotword_lengths"] = torch.tensor(hotword_lengths, dtype=torch.int32)
|
||||
if seaco_label_pad is not None:
|
||||
outputs["seaco_label_pad"] = seaco_label_pad
|
||||
return outputs
|
||||
@@ -0,0 +1,198 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
import logging
|
||||
import math
|
||||
import torch.distributed as dist
|
||||
from torch.utils.data import DistributedSampler
|
||||
from torch.utils.data import BatchSampler, Sampler
|
||||
import torch.distributed as dist
|
||||
import random
|
||||
from funasr.register import tables
|
||||
|
||||
|
||||
@tables.register("batch_sampler_classes", "EspnetStyleBatchSampler")
|
||||
def EspnetStyleBatchSampler_fn(dataset, **kwargs):
|
||||
"""Espnetstylebatchsampler fn.
|
||||
|
||||
Args:
|
||||
dataset: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
dataloader_args = {}
|
||||
|
||||
batch_sampler = EspnetStyleBatchSampler(dataset, **kwargs)
|
||||
dataloader_args["batch_sampler"] = batch_sampler
|
||||
dataloader_args["num_workers"] = kwargs.get("num_workers", 4)
|
||||
dataloader_args["pin_memory"] = kwargs.get("pin_memory", True)
|
||||
num_workers = dataloader_args.get("num_workers", 4)
|
||||
if num_workers > 0:
|
||||
dataloader_args["persistent_workers"] = kwargs.get("persistent_workers", True)
|
||||
dataloader_args["prefetch_factor"] = kwargs.get("prefetch_factor", 2)
|
||||
|
||||
return dataloader_args
|
||||
|
||||
|
||||
import torch
|
||||
from torch.utils.data import Dataset, DistributedSampler
|
||||
import math
|
||||
import random
|
||||
|
||||
|
||||
class EspnetStyleBatchSampler(DistributedSampler):
|
||||
def __init__(
|
||||
self,
|
||||
dataset,
|
||||
batch_size,
|
||||
batch_type="token",
|
||||
rank=None,
|
||||
num_replicas=None,
|
||||
rank_split=False,
|
||||
shuffle=True,
|
||||
drop_last=False,
|
||||
is_training: bool = True,
|
||||
sort_size: int = 1024,
|
||||
start_step: int = 0,
|
||||
**kwargs,
|
||||
):
|
||||
|
||||
"""Initialize EspnetStyleBatchSampler.
|
||||
|
||||
Args:
|
||||
dataset: TODO.
|
||||
batch_size: Number of samples per batch.
|
||||
batch_type: TODO.
|
||||
rank: TODO.
|
||||
num_replicas: TODO.
|
||||
rank_split: TODO.
|
||||
shuffle: TODO.
|
||||
drop_last: TODO.
|
||||
is_training: Boolean flag for training.
|
||||
sort_size: Size/dimension parameter.
|
||||
start_step: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
try:
|
||||
rank = dist.get_rank()
|
||||
num_replicas = dist.get_world_size()
|
||||
except:
|
||||
rank = 0
|
||||
num_replicas = 1
|
||||
# if rank_split:
|
||||
# logging.info(f"Warning, rank_split: {rank_split}, batch and shuffle data in local rank")
|
||||
# rank = 0
|
||||
# num_replicas = 1
|
||||
self.rank = rank
|
||||
self.num_replicas = num_replicas
|
||||
self.dataset = dataset
|
||||
self.batch_size = batch_size
|
||||
self.batch_type = batch_type
|
||||
self.is_training = is_training
|
||||
self.shuffle = shuffle and is_training
|
||||
self.drop_last = drop_last
|
||||
|
||||
self.total_size = len(self.dataset)
|
||||
self.num_samples = int(math.ceil(self.total_size / self.num_replicas))
|
||||
self.epoch = 0
|
||||
self.sort_size = sort_size * num_replicas
|
||||
self.max_token_length = kwargs.get("max_token_length", 2048)
|
||||
self.min_token_length = kwargs.get("min_token_length", 0)
|
||||
self.length_scale_source = kwargs.get("length_scale_source", 1.0)
|
||||
self.start_step = start_step
|
||||
self.batch_num = 1
|
||||
if self.start_step > 0:
|
||||
logging.info(f"Warning, start_step > 0, dataloader start from step: {self.start_step}")
|
||||
# super().__init__(dataset, num_replicas=num_replicas, rank=rank,
|
||||
# shuffle=shuffle, drop_last=drop_last)
|
||||
|
||||
def __iter__(self):
|
||||
"""Internal: iter ."""
|
||||
if self.shuffle:
|
||||
g = torch.Generator()
|
||||
g.manual_seed(self.epoch)
|
||||
random.seed(self.epoch)
|
||||
indices = torch.randperm(len(self.dataset), generator=g).tolist()
|
||||
else:
|
||||
indices = list(range(len(self.dataset)))
|
||||
|
||||
# Sort indices by sample length
|
||||
sorted_indices = sorted(indices, key=lambda idx: self.dataset.get_source_len(idx))
|
||||
|
||||
# Organize batches based on 'length' or 'example'
|
||||
buffer_batches = []
|
||||
batch = []
|
||||
max_len_in_batch = 0 # Tracks the max sample length within the current batch
|
||||
|
||||
for idx in sorted_indices:
|
||||
|
||||
# original_sample_length = self.dataset.get_source_len(idx)
|
||||
# if (
|
||||
# original_sample_length < self.min_token_length
|
||||
# or original_sample_length > self.max_token_length
|
||||
# ): # Skip samples that exceed the max length
|
||||
# continue
|
||||
|
||||
# sample_length = 1 if self.batch_type == "example" else original_sample_length
|
||||
|
||||
# Set sample_length based on the batch type
|
||||
if self.batch_type == "example":
|
||||
sample_length = 1
|
||||
elif self.batch_type == "token":
|
||||
sample_length = self.dataset.get_source_len(idx) + int(
|
||||
self.dataset.get_target_len(idx) * 1.2
|
||||
)
|
||||
else:
|
||||
sample_length = self.dataset.get_source_len(idx)
|
||||
# Calculate potential batch size with the new sample
|
||||
potential_batch_length = max(max_len_in_batch, sample_length) * (len(batch) + 1)
|
||||
# Add index to batch if it doesn't exceed batch size limit
|
||||
if potential_batch_length <= self.batch_size:
|
||||
batch.append(idx)
|
||||
max_len_in_batch = max(max_len_in_batch, sample_length)
|
||||
else:
|
||||
# Save the current batch and start a new one
|
||||
buffer_batches.append(batch)
|
||||
batch = [idx]
|
||||
max_len_in_batch = sample_length
|
||||
|
||||
# Add the last batch if it shouldn't be dropped
|
||||
if batch and (not self.drop_last or len(batch) * max_len_in_batch == self.batch_size):
|
||||
buffer_batches.append(batch)
|
||||
|
||||
# Shuffle the list of batches
|
||||
if self.shuffle:
|
||||
random.seed(self.epoch)
|
||||
random.shuffle(buffer_batches)
|
||||
|
||||
# Ensure each rank gets the same number of batches
|
||||
batches_per_rank = int(math.ceil(len(buffer_batches) / self.num_replicas))
|
||||
total_batches_needed = batches_per_rank * self.num_replicas
|
||||
extra_batches = total_batches_needed - len(buffer_batches)
|
||||
# Add extra batches by random selection, if needed
|
||||
buffer_batches += random.choices(buffer_batches, k=extra_batches)
|
||||
|
||||
# Allocate the batches to the current rank
|
||||
start_idx = self.rank * batches_per_rank
|
||||
end_idx = start_idx + batches_per_rank
|
||||
rank_batches = buffer_batches[start_idx + self.start_step : end_idx]
|
||||
|
||||
self.batch_num = len(rank_batches)
|
||||
|
||||
logging.info(
|
||||
f"rank: {self.rank}, dataloader start from step: {self.start_step}, batch_num: {end_idx-start_idx}, batch_num_after_step: {len(rank_batches)}"
|
||||
)
|
||||
# Return an iterator over the batches for the current rank
|
||||
return iter(rank_batches)
|
||||
|
||||
def __len__(self):
|
||||
# Calculate the number of batches per epoch for the current rank
|
||||
"""Internal: len ."""
|
||||
return self.batch_num
|
||||
|
||||
def set_epoch(self, epoch):
|
||||
# Set the epoch for shuffling
|
||||
"""Set epoch.
|
||||
|
||||
Args:
|
||||
epoch: TODO.
|
||||
"""
|
||||
self.epoch = epoch
|
||||
@@ -0,0 +1,173 @@
|
||||
import os
|
||||
import json
|
||||
import torch
|
||||
import logging
|
||||
|
||||
import librosa
|
||||
import random
|
||||
import torch.distributed as dist
|
||||
|
||||
from funasr.register import tables
|
||||
|
||||
|
||||
@tables.register("index_ds_classes", "IndexDSJsonl")
|
||||
@tables.register("index_ds_classes", "IndexDSJsonlRankFull")
|
||||
@tables.register("index_ds_classes", "IndexDSJsonlRankSplit")
|
||||
class IndexDSJsonlRankFull(torch.utils.data.Dataset):
|
||||
|
||||
def __init__(self, path: str, **kwargs):
|
||||
"""Initialize IndexDSJsonlRankFull.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
self.max_source_length = kwargs.get("max_source_length", 2048)
|
||||
self.min_source_length = kwargs.get("min_source_length", 0)
|
||||
self.max_target_length = kwargs.get("max_target_length", 2048)
|
||||
self.min_target_length = kwargs.get("min_target_length", 0)
|
||||
self.max_token_length = kwargs.get("max_token_length", 2200)
|
||||
|
||||
is_training = kwargs.get("is_training", True)
|
||||
if not (path.endswith(".jsonl") or path.endswith(".json")):
|
||||
# jsonl list file
|
||||
data_split_num = kwargs.get("data_split_num", 1)
|
||||
data_split_i = kwargs.get("data_split_i", 0)
|
||||
|
||||
if not is_training:
|
||||
data_split_num = 1
|
||||
data_split_i = 0
|
||||
with open(path, encoding="utf-8") as fin:
|
||||
file_list_all = fin.readlines()
|
||||
|
||||
num_per_slice = (len(file_list_all) - 1) // data_split_num + 1 # 16
|
||||
file_list = file_list_all[
|
||||
data_split_i * num_per_slice : (data_split_i + 1) * num_per_slice
|
||||
]
|
||||
logging.info(
|
||||
f"is_training: {is_training}, data_split_num: {data_split_num}, data_split_i: {data_split_i}, \nfile_list: {file_list}, \nfile_list_all: {file_list_all}"
|
||||
)
|
||||
|
||||
else:
|
||||
file_list = [path]
|
||||
|
||||
# total_num = len(file_list)
|
||||
# try:
|
||||
# rank = dist.get_rank()
|
||||
# world_size = dist.get_world_size()
|
||||
# except:
|
||||
# rank = 0
|
||||
# world_size = 1
|
||||
# logging.info("distributed is not initialized, only single shard")
|
||||
#
|
||||
# if not kwargs.get("rank_split", False):
|
||||
# logging.info(f"Warning, rank_split disenabled, batch and shuffle data in global")
|
||||
# rank = 0
|
||||
# world_size = 1
|
||||
#
|
||||
# num_per_rank = total_num // world_size
|
||||
# if num_per_rank * world_size < total_num:
|
||||
# logging.info(f"Warning, jsonl file:{total_num} could not be divided by world_size: {world_size}, {path}")
|
||||
# total_num_needed = num_per_rank * world_size
|
||||
#
|
||||
# extra_num = total_num_needed - total_num
|
||||
# file_list_tmp = random.choices(file_list, k=extra_num)
|
||||
# file_list += file_list_tmp
|
||||
# logging.info(f"Warning, after random choices: {file_list}")
|
||||
#
|
||||
# file_list_rank = file_list[rank * num_per_rank:(rank + 1) * num_per_rank]
|
||||
#
|
||||
# logging.info(
|
||||
# f"is_training: {is_training}, file_list_rank: {file_list_rank}")
|
||||
|
||||
# contents = []
|
||||
# for file_json in file_list_rank:
|
||||
contents = []
|
||||
for file_json in file_list:
|
||||
with open(file_json.strip(), encoding="utf-8") as fin:
|
||||
for line in fin:
|
||||
data = json.loads(line.strip())
|
||||
if "text" in data: # for sft
|
||||
contents.append(data["text"])
|
||||
if "source" in data: # for speech lab pretrain
|
||||
prompt = data.get("prompt", "<ASR>")
|
||||
source = data["source"].replace(
|
||||
"/cpfs01", "/cpfs_speech/data"
|
||||
) # only use in alibaba gpu group: .replace("/cpfs01", "/cpfs_speech/data")
|
||||
target = data["target"]
|
||||
source_len = data.get("source_len", 1)
|
||||
target_len = data.get("target_len", 0)
|
||||
text_language = data.get("text_language", "")
|
||||
if "aishell" in source and text_language != "en":
|
||||
target = target.replace(" ", "")
|
||||
if (
|
||||
source_len < self.min_source_length
|
||||
or source_len > self.max_source_length
|
||||
):
|
||||
continue
|
||||
if (
|
||||
target_len < self.min_target_length
|
||||
or target_len > self.max_target_length
|
||||
):
|
||||
continue
|
||||
|
||||
if (source_len + target_len) > self.max_token_length:
|
||||
continue
|
||||
|
||||
contents_i = {
|
||||
"source": source,
|
||||
"prompt": prompt,
|
||||
"target": target,
|
||||
"source_len": source_len,
|
||||
"target_len": target_len,
|
||||
}
|
||||
text_language = data.get("text_language", None)
|
||||
if text_language is not None:
|
||||
contents_i["text_language"] = text_language
|
||||
if "emo_target" in data:
|
||||
contents_i["emo_target"] = data["emo_target"]
|
||||
if "event_target" in data:
|
||||
contents_i["event_target"] = data["event_target"]
|
||||
if "with_or_wo_itn" in data:
|
||||
contents_i["with_or_wo_itn"] = data["with_or_wo_itn"]
|
||||
# audio_language = data.get("audio_language", None)
|
||||
# if audio_language is not None:
|
||||
# contents_i["audio_language"] = audio_language
|
||||
contents.append(contents_i)
|
||||
|
||||
self.contents = contents
|
||||
|
||||
logging.info("total_num of samplers: {}, {}".format(len(self.contents), path))
|
||||
|
||||
def __len__(self):
|
||||
"""Internal: len ."""
|
||||
return len(self.contents)
|
||||
|
||||
def __getitem__(self, index):
|
||||
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
data = self.contents[index]
|
||||
|
||||
return data
|
||||
|
||||
def get_source_len(self, data_dict):
|
||||
"""Get source len.
|
||||
|
||||
Args:
|
||||
data_dict: TODO.
|
||||
"""
|
||||
return data_dict.get("source_len", 1)
|
||||
|
||||
def get_target_len(self, data_dict):
|
||||
|
||||
"""Get target len.
|
||||
|
||||
Args:
|
||||
data_dict: TODO.
|
||||
"""
|
||||
return data_dict.get("target_len", 0)
|
||||
@@ -0,0 +1,77 @@
|
||||
import os
|
||||
import json
|
||||
import torch
|
||||
import logging
|
||||
import hydra
|
||||
from omegaconf import DictConfig, OmegaConf
|
||||
import concurrent.futures
|
||||
import librosa
|
||||
import torch.distributed as dist
|
||||
|
||||
|
||||
def gen_scp_from_jsonl(jsonl_file, data_type_list, wav_scp_file, text_file):
|
||||
|
||||
"""Gen scp from jsonl.
|
||||
|
||||
Args:
|
||||
jsonl_file: TODO.
|
||||
data_type_list: TODO.
|
||||
wav_scp_file: TODO.
|
||||
text_file: TODO.
|
||||
"""
|
||||
wav_f = open(wav_scp_file, "w")
|
||||
text_f = open(text_file, "w")
|
||||
with open(jsonl_file, encoding="utf-8") as fin:
|
||||
for line in fin:
|
||||
data = json.loads(line.strip())
|
||||
|
||||
prompt = data.get("prompt", "<ASR>")
|
||||
source = data[data_type_list[0]]
|
||||
target = data[data_type_list[1]]
|
||||
source_len = data.get("source_len", 1)
|
||||
target_len = data.get("target_len", 0)
|
||||
if "aishell" in source:
|
||||
target = target.replace(" ", "")
|
||||
key = data["key"]
|
||||
wav_f.write(f"{key}\t{source}\n")
|
||||
wav_f.flush()
|
||||
text_f.write(f"{key}\t{target}\n")
|
||||
text_f.flush()
|
||||
|
||||
wav_f.close()
|
||||
text_f.close()
|
||||
|
||||
|
||||
@hydra.main(config_name=None, version_base=None)
|
||||
def main_hydra(cfg: DictConfig):
|
||||
|
||||
"""Main hydra.
|
||||
|
||||
Args:
|
||||
cfg: Configuration overrides.
|
||||
"""
|
||||
kwargs = OmegaConf.to_container(cfg, resolve=True)
|
||||
print(kwargs)
|
||||
|
||||
scp_file_list = kwargs.get(
|
||||
"scp_file_list",
|
||||
("/Users/zhifu/funasr1.0/test_local/wav.scp", "/Users/zhifu/funasr1.0/test_local/text.txt"),
|
||||
)
|
||||
if isinstance(scp_file_list, str):
|
||||
scp_file_list = eval(scp_file_list)
|
||||
data_type_list = kwargs.get("data_type_list", ("source", "target"))
|
||||
jsonl_file = kwargs.get(
|
||||
"jsonl_file_in", "/Users/zhifu/funasr1.0/test_local/audio_datasets.jsonl"
|
||||
)
|
||||
gen_scp_from_jsonl(jsonl_file, data_type_list, *scp_file_list)
|
||||
|
||||
|
||||
"""
|
||||
python -m funasr.datasets.audio_datasets.json2scp \
|
||||
++scp_file_list='["/Users/zhifu/funasr1.0/test_local/wav.scp", "/Users/zhifu/funasr1.0/test_local/text.txt"]' \
|
||||
++data_type_list='["source", "target"]' \
|
||||
++jsonl_file_in=/Users/zhifu/funasr1.0/test_local/audio_datasets.jsonl
|
||||
"""
|
||||
|
||||
if __name__ == "__main__":
|
||||
main_hydra()
|
||||
@@ -0,0 +1,82 @@
|
||||
import os
|
||||
import json
|
||||
import torch
|
||||
import logging
|
||||
import concurrent.futures
|
||||
import librosa
|
||||
import torch.distributed as dist
|
||||
from typing import Collection
|
||||
import torch
|
||||
import torchaudio
|
||||
from torch import nn
|
||||
import random
|
||||
import re
|
||||
from funasr.tokenizer.cleaner import TextCleaner
|
||||
from funasr.register import tables
|
||||
|
||||
|
||||
@tables.register("preprocessor_classes", "SpeechPreprocessSpeedPerturb")
|
||||
class SpeechPreprocessSpeedPerturb(nn.Module):
|
||||
def __init__(self, speed_perturb: list = None, **kwargs):
|
||||
"""Initialize SpeechPreprocessSpeedPerturb.
|
||||
|
||||
Args:
|
||||
speed_perturb: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
self.speed_perturb = speed_perturb
|
||||
|
||||
def forward(self, waveform, fs, **kwargs):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
waveform: TODO.
|
||||
fs: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
if self.speed_perturb is None:
|
||||
return waveform
|
||||
speed = random.choice(self.speed_perturb)
|
||||
if speed != 1.0:
|
||||
if not isinstance(waveform, torch.Tensor):
|
||||
waveform = torch.tensor(waveform)
|
||||
waveform, _ = torchaudio.sox_effects.apply_effects_tensor(
|
||||
waveform.view(1, -1), fs, [["speed", str(speed)], ["rate", str(fs)]]
|
||||
)
|
||||
waveform = waveform.view(-1)
|
||||
|
||||
return waveform
|
||||
|
||||
|
||||
@tables.register("preprocessor_classes", "TextPreprocessSegDict")
|
||||
class TextPreprocessSegDict(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
seg_dict: str = None,
|
||||
text_cleaner: Collection[str] = None,
|
||||
split_with_space: bool = False,
|
||||
**kwargs
|
||||
):
|
||||
"""Initialize TextPreprocessSegDict.
|
||||
|
||||
Args:
|
||||
seg_dict: TODO.
|
||||
text_cleaner: TODO.
|
||||
split_with_space: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
self.text_cleaner = TextCleaner(text_cleaner)
|
||||
|
||||
def forward(self, text, **kwargs):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
text: Text tensor or string input.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
text = self.text_cleaner(text)
|
||||
|
||||
return text
|
||||
@@ -0,0 +1,592 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
import logging
|
||||
import math
|
||||
import random
|
||||
import torch.distributed as dist
|
||||
from torch.utils.data import DistributedSampler
|
||||
from torch.utils.data import BatchSampler, Sampler
|
||||
import torch.distributed as dist
|
||||
|
||||
from funasr.register import tables
|
||||
|
||||
|
||||
@tables.register("batch_sampler_classes", "BatchSampler")
|
||||
@tables.register("batch_sampler_classes", "CustomDistributedBatchSampler")
|
||||
@tables.register("batch_sampler_classes", "CustomDistributedDynamicBatchSampler")
|
||||
@tables.register("batch_sampler_classes", "DynamicBatchLocalShuffleSampler")
|
||||
@tables.register("batch_sampler_classes", "RankFullLocalShuffleBatchSampler")
|
||||
@tables.register("batch_sampler_classes", "RankFullLocalShuffleDynamicBatchSampler")
|
||||
def CustomDistributedBatchSampler_fn(dataset, **kwargs):
|
||||
"""Customdistributedbatchsampler fn.
|
||||
|
||||
Args:
|
||||
dataset: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
dataloader_args = {}
|
||||
batch_type = kwargs.get("batch_type", "example")
|
||||
if batch_type == "example":
|
||||
batch_sampler = CustomDistributedBatchSampler(dataset, **kwargs)
|
||||
|
||||
else:
|
||||
if kwargs.get("sort_size", -1) > 0:
|
||||
batch_sampler = CustomDistributedBufferDynamicBatchSampler(dataset, **kwargs)
|
||||
else:
|
||||
batch_sampler = CustomDistributedDynamicBatchSampler(dataset, **kwargs)
|
||||
# batch_sampler = CustomDistributedDynamicBatchSampler(dataset, **kwargs)
|
||||
|
||||
dataloader_args["batch_sampler"] = batch_sampler
|
||||
dataloader_args["num_workers"] = kwargs.get("num_workers", 4)
|
||||
dataloader_args["pin_memory"] = kwargs.get("pin_memory", True)
|
||||
num_workers = dataloader_args.get("num_workers", 4)
|
||||
if num_workers > 0:
|
||||
dataloader_args["persistent_workers"] = kwargs.get("persistent_workers", True)
|
||||
dataloader_args["prefetch_factor"] = kwargs.get("prefetch_factor", 2)
|
||||
|
||||
return dataloader_args
|
||||
|
||||
|
||||
class CustomDistributedBatchSampler(Sampler):
|
||||
def __init__(
|
||||
self,
|
||||
dataset,
|
||||
batch_size,
|
||||
num_replicas=None,
|
||||
rank=None,
|
||||
shuffle=True,
|
||||
drop_last=False,
|
||||
is_training: bool = True,
|
||||
**kwargs,
|
||||
):
|
||||
|
||||
"""Initialize CustomDistributedBatchSampler.
|
||||
|
||||
Args:
|
||||
dataset: TODO.
|
||||
batch_size: Number of samples per batch.
|
||||
num_replicas: TODO.
|
||||
rank: TODO.
|
||||
shuffle: TODO.
|
||||
drop_last: TODO.
|
||||
is_training: Boolean flag for training.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
try:
|
||||
rank = dist.get_rank()
|
||||
num_replicas = dist.get_world_size()
|
||||
except:
|
||||
rank = 0
|
||||
num_replicas = 1
|
||||
self.rank = rank
|
||||
self.num_replicas = num_replicas
|
||||
self.dataset = dataset
|
||||
self.batch_size = batch_size
|
||||
self.is_training = is_training
|
||||
self.shuffle = shuffle and is_training
|
||||
self.drop_last = drop_last
|
||||
# self.total_size = len(dataset)
|
||||
if self.drop_last:
|
||||
self.total_size = (len(self.dataset) // (batch_size * num_replicas)) * (
|
||||
batch_size * num_replicas
|
||||
)
|
||||
else:
|
||||
self.total_size = math.ceil(len(self.dataset) / (batch_size * num_replicas)) * (
|
||||
batch_size * num_replicas
|
||||
)
|
||||
self.num_samples = int(self.total_size // self.num_replicas)
|
||||
self.epoch = 0
|
||||
self.max_token_length = kwargs.get("max_token_length", None)
|
||||
self.length_scale_source = kwargs.get("length_scale_source", 1.0)
|
||||
|
||||
def __iter__(self):
|
||||
# Generate a list of indices
|
||||
"""Internal: iter ."""
|
||||
if self.shuffle:
|
||||
g = torch.Generator()
|
||||
g.manual_seed(self.epoch)
|
||||
indices = torch.randperm(len(self.dataset), generator=g).tolist()
|
||||
else:
|
||||
indices = list(range(len(self.dataset)))
|
||||
|
||||
# Add extra samples to make it evenly divisible
|
||||
padding_size = self.total_size - len(indices)
|
||||
if padding_size <= len(indices):
|
||||
indices += indices[:padding_size]
|
||||
else:
|
||||
indices += (
|
||||
indices * (padding_size // len(indices)) + indices[: padding_size % len(indices)]
|
||||
)
|
||||
|
||||
assert len(indices) == self.total_size
|
||||
|
||||
# Subsample
|
||||
indices = indices[self.rank : self.total_size : self.num_replicas]
|
||||
assert len(indices) == self.num_samples
|
||||
|
||||
# Filter out indices with length greater than the max length, if provided
|
||||
if self.max_token_length is not None:
|
||||
filtered_indices = []
|
||||
for idx in indices:
|
||||
source_len = self.dataset.get_source_len(idx) / self.length_scale_source
|
||||
if source_len <= self.max_token_length:
|
||||
filtered_indices.append(idx)
|
||||
indices = filtered_indices
|
||||
|
||||
# Now that we have only the indices for this replica, chunk them into batches
|
||||
batches = [
|
||||
indices[i : i + self.batch_size] for i in range(0, len(indices), self.batch_size)
|
||||
]
|
||||
|
||||
# Drop the last batch if it's not full and drop_last is True
|
||||
if self.drop_last and len(batches[-1]) != self.batch_size:
|
||||
batches = batches[:-1]
|
||||
|
||||
return iter(batches)
|
||||
|
||||
def __len__(self):
|
||||
|
||||
"""Internal: len ."""
|
||||
return self.num_samples // self.batch_size
|
||||
|
||||
def set_epoch(self, epoch):
|
||||
"""Set epoch.
|
||||
|
||||
Args:
|
||||
epoch: TODO.
|
||||
"""
|
||||
self.epoch = epoch
|
||||
|
||||
|
||||
class CustomDistributedBufferBatchSampler(Sampler):
|
||||
def __init__(
|
||||
self,
|
||||
dataset,
|
||||
batch_size,
|
||||
num_replicas=None,
|
||||
rank=None,
|
||||
shuffle=True,
|
||||
drop_last=False,
|
||||
is_training: bool = True,
|
||||
sort_size: int = 1024,
|
||||
**kwargs,
|
||||
):
|
||||
|
||||
"""Initialize CustomDistributedBufferBatchSampler.
|
||||
|
||||
Args:
|
||||
dataset: TODO.
|
||||
batch_size: Number of samples per batch.
|
||||
num_replicas: TODO.
|
||||
rank: TODO.
|
||||
shuffle: TODO.
|
||||
drop_last: TODO.
|
||||
is_training: Boolean flag for training.
|
||||
sort_size: Size/dimension parameter.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
try:
|
||||
rank = dist.get_rank()
|
||||
num_replicas = dist.get_world_size()
|
||||
except:
|
||||
rank = 0
|
||||
num_replicas = 1
|
||||
self.rank = rank
|
||||
self.num_replicas = num_replicas
|
||||
self.dataset = dataset
|
||||
self.batch_size = batch_size
|
||||
self.is_training = is_training
|
||||
self.shuffle = shuffle and is_training
|
||||
self.drop_last = drop_last
|
||||
# self.total_size = len(dataset)
|
||||
if self.drop_last:
|
||||
self.total_size = (len(self.dataset) // (batch_size * num_replicas)) * (
|
||||
batch_size * num_replicas
|
||||
)
|
||||
else:
|
||||
self.total_size = math.ceil(len(self.dataset) / (batch_size * num_replicas)) * (
|
||||
batch_size * num_replicas
|
||||
)
|
||||
self.num_samples = int(self.total_size // self.num_replicas)
|
||||
self.epoch = 0
|
||||
self.max_token_length = kwargs.get("max_token_length", None)
|
||||
self.length_scale_source = kwargs.get("length_scale_source", 1.0)
|
||||
self.sort_size = sort_size
|
||||
|
||||
def __iter__(self):
|
||||
# Generate a list of indices
|
||||
"""Internal: iter ."""
|
||||
if self.shuffle:
|
||||
g = torch.Generator()
|
||||
g.manual_seed(self.epoch)
|
||||
indices = torch.randperm(len(self.dataset), generator=g).tolist()
|
||||
else:
|
||||
indices = list(range(len(self.dataset)))
|
||||
|
||||
# Add extra samples to make it evenly divisible
|
||||
padding_size = self.total_size - len(indices)
|
||||
if padding_size <= len(indices):
|
||||
indices += indices[:padding_size]
|
||||
else:
|
||||
indices += (
|
||||
indices * (padding_size // len(indices)) + indices[: padding_size % len(indices)]
|
||||
)
|
||||
|
||||
assert len(indices) == self.total_size
|
||||
|
||||
# Subsample
|
||||
indices = indices[self.rank : self.total_size : self.num_replicas]
|
||||
assert len(indices) == self.num_samples
|
||||
|
||||
# Filter out indices with length greater than the max length, if provided
|
||||
if self.max_token_length is not None:
|
||||
filtered_indices = []
|
||||
for idx in indices:
|
||||
source_len = self.dataset.get_source_len(idx) / self.length_scale_source
|
||||
if source_len <= self.max_token_length:
|
||||
filtered_indices.append(idx)
|
||||
indices = filtered_indices
|
||||
|
||||
# Buffer sorting logic
|
||||
sorted_batches = []
|
||||
buffer = []
|
||||
|
||||
for idx in indices:
|
||||
buffer.append(idx)
|
||||
if len(buffer) >= self.sort_size:
|
||||
# Sort the buffer based on some criteria, e.g., dataset sample length
|
||||
buffer.sort(key=lambda x: self.dataset.get_source_len(x))
|
||||
sorted_batches.extend(self._create_batches_from_buffer(buffer))
|
||||
buffer = []
|
||||
|
||||
# Handle the remaining items in the buffer
|
||||
if buffer:
|
||||
buffer.sort(key=lambda x: self.dataset.get_source_len(x))
|
||||
sorted_batches.extend(self._create_batches_from_buffer(buffer))
|
||||
|
||||
return iter(sorted_batches)
|
||||
|
||||
def _create_batches_from_buffer(self, buffer):
|
||||
# Function to convert the sorted buffer into batches
|
||||
"""Internal: create batches from buffer.
|
||||
|
||||
Args:
|
||||
buffer: TODO.
|
||||
"""
|
||||
batched_buffer = [
|
||||
buffer[i : i + self.batch_size] for i in range(0, len(buffer), self.batch_size)
|
||||
]
|
||||
if self.drop_last and len(batched_buffer[-1]) != self.batch_size:
|
||||
batched_buffer = batched_buffer[:-1]
|
||||
return batched_buffer
|
||||
|
||||
def __len__(self):
|
||||
|
||||
"""Internal: len ."""
|
||||
return self.num_samples // self.batch_size
|
||||
|
||||
def set_epoch(self, epoch):
|
||||
"""Set epoch.
|
||||
|
||||
Args:
|
||||
epoch: TODO.
|
||||
"""
|
||||
self.epoch = epoch
|
||||
|
||||
|
||||
class CustomDistributedDynamicBatchSampler(DistributedSampler):
|
||||
def __init__(
|
||||
self,
|
||||
dataset,
|
||||
batch_size,
|
||||
num_replicas=None,
|
||||
rank=None,
|
||||
shuffle=True,
|
||||
drop_last=False,
|
||||
is_training: bool = True,
|
||||
**kwargs,
|
||||
):
|
||||
|
||||
"""Initialize CustomDistributedDynamicBatchSampler.
|
||||
|
||||
Args:
|
||||
dataset: TODO.
|
||||
batch_size: Number of samples per batch.
|
||||
num_replicas: TODO.
|
||||
rank: TODO.
|
||||
shuffle: TODO.
|
||||
drop_last: TODO.
|
||||
is_training: Boolean flag for training.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
try:
|
||||
rank = dist.get_rank()
|
||||
num_replicas = dist.get_world_size()
|
||||
except:
|
||||
rank = 0
|
||||
num_replicas = 1
|
||||
self.rank = rank
|
||||
self.num_replicas = num_replicas
|
||||
self.dataset = dataset
|
||||
self.batch_size = batch_size
|
||||
self.is_training = is_training
|
||||
self.shuffle = shuffle and is_training
|
||||
self.drop_last = drop_last
|
||||
|
||||
self.total_size = len(self.dataset)
|
||||
# self.num_samples = int(math.ceil(self.total_size / self.num_replicas))
|
||||
self.epoch = 0
|
||||
self.max_token_length = kwargs.get("max_token_length", 2048)
|
||||
self.length_scale_source = kwargs.get("length_scale_source", 1.0)
|
||||
|
||||
def __iter__(self):
|
||||
"""Internal: iter ."""
|
||||
if self.shuffle:
|
||||
g = torch.Generator()
|
||||
g.manual_seed(self.epoch)
|
||||
indices = torch.randperm(len(self.dataset), generator=g).tolist()
|
||||
else:
|
||||
indices = list(range(len(self.dataset)))
|
||||
|
||||
indices = indices[self.rank : self.total_size : self.num_replicas]
|
||||
|
||||
batches = []
|
||||
batch = []
|
||||
max_len_in_batch = 0
|
||||
current_batch_length = 0
|
||||
|
||||
for idx in indices:
|
||||
sample_length = self.dataset.get_source_len(idx)
|
||||
if sample_length > self.max_token_length:
|
||||
continue
|
||||
potential_batch_length = (
|
||||
max_len_in_batch if sample_length < max_len_in_batch else sample_length
|
||||
) * (len(batch) + 1)
|
||||
|
||||
if potential_batch_length <= self.batch_size:
|
||||
batch.append(idx)
|
||||
if sample_length > max_len_in_batch:
|
||||
max_len_in_batch = sample_length
|
||||
# current_batch_length = max_len_in_batch * len(batch)
|
||||
else:
|
||||
batches.append(batch)
|
||||
batch = [idx]
|
||||
max_len_in_batch = sample_length
|
||||
# current_batch_length = max_len_in_batch
|
||||
|
||||
# Add the last batch if it's not empty and we're not dropping it
|
||||
if batch and (not self.drop_last or len(batch) * max_len_in_batch == self.batch_size):
|
||||
batches.append(batch)
|
||||
|
||||
return iter(batches)
|
||||
|
||||
def __len__(self):
|
||||
|
||||
"""Internal: len ."""
|
||||
return 1
|
||||
|
||||
def set_epoch(self, epoch):
|
||||
"""Set epoch.
|
||||
|
||||
Args:
|
||||
epoch: TODO.
|
||||
"""
|
||||
self.epoch = epoch
|
||||
|
||||
|
||||
class CustomDistributedBufferDynamicBatchSampler(DistributedSampler):
|
||||
def __init__(
|
||||
self,
|
||||
dataset,
|
||||
batch_size,
|
||||
batch_type="token",
|
||||
num_replicas=None,
|
||||
rank=None,
|
||||
rank_split=False,
|
||||
shuffle=True,
|
||||
drop_last=False,
|
||||
is_training: bool = True,
|
||||
sort_size: int = 1024,
|
||||
start_step: int = 0,
|
||||
**kwargs,
|
||||
):
|
||||
|
||||
"""Initialize CustomDistributedBufferDynamicBatchSampler.
|
||||
|
||||
Args:
|
||||
dataset: TODO.
|
||||
batch_size: Number of samples per batch.
|
||||
batch_type: TODO.
|
||||
num_replicas: TODO.
|
||||
rank: TODO.
|
||||
rank_split: TODO.
|
||||
shuffle: TODO.
|
||||
drop_last: TODO.
|
||||
is_training: Boolean flag for training.
|
||||
sort_size: Size/dimension parameter.
|
||||
start_step: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
try:
|
||||
rank = dist.get_rank()
|
||||
num_replicas = dist.get_world_size()
|
||||
except:
|
||||
rank = 0
|
||||
num_replicas = 1
|
||||
|
||||
# if rank_split:
|
||||
# logging.info(f"Warning, rank_split: {rank_split}, batch and shuffle data in local rank")
|
||||
# rank = 0
|
||||
# num_replicas = 1
|
||||
|
||||
self.rank = rank
|
||||
self.num_replicas = num_replicas
|
||||
self.dataset = dataset
|
||||
self.batch_size = batch_size
|
||||
self.batch_type = batch_type
|
||||
self.is_training = is_training
|
||||
self.shuffle = shuffle and is_training
|
||||
self.drop_last = drop_last
|
||||
|
||||
self.total_size = len(self.dataset)
|
||||
self.num_samples = int(math.ceil(self.total_size / self.num_replicas))
|
||||
self.epoch = 0
|
||||
self.sort_size = sort_size * num_replicas
|
||||
self.max_token_length = kwargs.get("max_token_length", 2048)
|
||||
self.length_scale_source = kwargs.get("length_scale_source", 1.0)
|
||||
self.batch_size_sample_max = kwargs.get("batch_size_sample_max", 200)
|
||||
self.start_step = start_step
|
||||
self.batch_num = 1
|
||||
if self.start_step > 0:
|
||||
logging.info(f"Warning, start_step > 0, dataloader start from step: {self.start_step}")
|
||||
# super().__init__(
|
||||
# dataset, num_replicas=num_replicas, rank=rank, shuffle=shuffle, drop_last=drop_last
|
||||
# )
|
||||
|
||||
def __iter__(self):
|
||||
"""Internal: iter ."""
|
||||
if self.shuffle:
|
||||
g = torch.Generator()
|
||||
g.manual_seed(self.epoch)
|
||||
random.seed(self.epoch)
|
||||
|
||||
indices = torch.randperm(len(self.dataset), generator=g).tolist()
|
||||
else:
|
||||
indices = list(range(len(self.dataset)))
|
||||
|
||||
# Create sorted buffers and form batches
|
||||
buffer_batches = []
|
||||
for i in range(0, len(indices), self.sort_size):
|
||||
buffer = sorted(
|
||||
indices[i : i + self.sort_size], key=lambda idx: self.dataset.get_source_len(idx)
|
||||
)
|
||||
batch = []
|
||||
max_len_in_batch = 0
|
||||
count = 1
|
||||
for idx in buffer:
|
||||
original_sample_length = self.dataset.get_source_len(idx)
|
||||
if original_sample_length > self.max_token_length:
|
||||
continue
|
||||
sample_length = 1 if self.batch_type == "example" else original_sample_length
|
||||
potential_batch_length = max(max_len_in_batch, sample_length) * (len(batch) + 1)
|
||||
if potential_batch_length <= self.batch_size and count < self.batch_size_sample_max:
|
||||
batch.append(idx)
|
||||
max_len_in_batch = max(max_len_in_batch, sample_length)
|
||||
count += 1
|
||||
else:
|
||||
buffer_batches.append(batch)
|
||||
batch = [idx]
|
||||
max_len_in_batch = sample_length
|
||||
count = 1
|
||||
if batch:
|
||||
buffer_batches.append(batch)
|
||||
|
||||
# Ensure each rank gets the same number of batches, duplicate data if needed
|
||||
batches_per_rank = math.ceil(len(buffer_batches) / self.num_replicas)
|
||||
total_batches_needed = batches_per_rank * self.num_replicas
|
||||
|
||||
extra_batches = total_batches_needed - len(buffer_batches)
|
||||
buffer_batches += random.choices(buffer_batches, k=extra_batches)
|
||||
|
||||
# Evenly distribute batches from buffer_batches to each rank
|
||||
rank_batches = [[] for _ in range(self.num_replicas)]
|
||||
for i, batch in enumerate(buffer_batches):
|
||||
rank_batches[i % self.num_replicas].append(batch)
|
||||
|
||||
# Assign all batches for the current rank directly
|
||||
final_batches = rank_batches[self.rank][self.start_step :]
|
||||
self.batch_num = len(final_batches)
|
||||
|
||||
logging.info(
|
||||
f"rank: {self.rank}, dataloader start from step: {self.start_step}, batch_num: {len(rank_batches[self.rank])}, after: {self.batch_num}"
|
||||
)
|
||||
return iter(final_batches)
|
||||
|
||||
def __len__(self):
|
||||
# Calculate the number of batches per epoch for the current rank
|
||||
"""Internal: len ."""
|
||||
return self.batch_num
|
||||
|
||||
def set_epoch(self, epoch):
|
||||
"""Set epoch.
|
||||
|
||||
Args:
|
||||
epoch: TODO.
|
||||
"""
|
||||
self.epoch = epoch
|
||||
|
||||
|
||||
class DistributedSamplerWarp(BatchSampler):
|
||||
def __init__(
|
||||
self, dataset, batch_size, num_replicas=None, rank=None, shuffle=True, drop_last=False
|
||||
):
|
||||
"""Initialize DistributedSamplerWarp.
|
||||
|
||||
Args:
|
||||
dataset: TODO.
|
||||
batch_size: Number of samples per batch.
|
||||
num_replicas: TODO.
|
||||
rank: TODO.
|
||||
shuffle: TODO.
|
||||
drop_last: TODO.
|
||||
"""
|
||||
if num_replicas is None:
|
||||
if not torch.distributed.is_available():
|
||||
raise RuntimeError("Requires distributed package to be available")
|
||||
num_replicas = torch.distributed.get_world_size()
|
||||
if rank is None:
|
||||
if not torch.distributed.is_available():
|
||||
raise RuntimeError("Requires distributed package to be available")
|
||||
rank = torch.distributed.get_rank()
|
||||
|
||||
self.dataset = dataset
|
||||
self.batch_size = batch_size
|
||||
self.num_replicas = num_replicas
|
||||
self.rank = rank
|
||||
self.shuffle = shuffle
|
||||
self.drop_last = drop_last
|
||||
|
||||
# Create an instance of the DistributedSampler
|
||||
self.sampler = DistributedSampler(
|
||||
self.dataset, num_replicas=self.num_replicas, rank=self.rank, shuffle=self.shuffle
|
||||
)
|
||||
|
||||
# Call BatchSampler's constructor
|
||||
super().__init__(self.sampler, batch_size, drop_last)
|
||||
|
||||
def __iter__(self):
|
||||
# If we shuffle, we need to call the set_epoch method
|
||||
"""Internal: iter ."""
|
||||
if self.shuffle:
|
||||
self.sampler.set_epoch(self.epoch)
|
||||
|
||||
# Generate batch indices using the parent class
|
||||
return super().__iter__()
|
||||
|
||||
def set_epoch(self, epoch):
|
||||
"""Set epoch.
|
||||
|
||||
Args:
|
||||
epoch: TODO.
|
||||
"""
|
||||
self.epoch = epoch
|
||||
@@ -0,0 +1,148 @@
|
||||
import os
|
||||
import json
|
||||
import torch
|
||||
import logging
|
||||
import hydra
|
||||
from omegaconf import DictConfig, OmegaConf
|
||||
import concurrent.futures
|
||||
import librosa
|
||||
import torch.distributed as dist
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
def gen_jsonl_from_wav_text_list(
|
||||
path, data_type_list=("source", "target"), jsonl_file_out: str = None, **kwargs
|
||||
):
|
||||
"""Gen jsonl from wav text list.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
data_type_list: TODO.
|
||||
jsonl_file_out: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
try:
|
||||
rank = dist.get_rank()
|
||||
world_size = dist.get_world_size()
|
||||
except:
|
||||
rank = 0
|
||||
world_size = 1
|
||||
|
||||
cpu_cores = os.cpu_count() or 1
|
||||
print(f"convert wav.scp text to jsonl, ncpu: {cpu_cores}")
|
||||
if rank == 0:
|
||||
json_dict = {}
|
||||
for data_type, data_file in zip(data_type_list, path):
|
||||
json_dict[data_type] = {}
|
||||
with open(data_file, "r") as f:
|
||||
|
||||
data_file_lists = f.readlines()
|
||||
lines_for_each_th = (len(data_file_lists) - 1) // cpu_cores + 1
|
||||
task_num = cpu_cores if len(data_file_lists) > cpu_cores else 1
|
||||
# import pdb;pdb.set_trace()
|
||||
if task_num > 1:
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=cpu_cores) as executor:
|
||||
|
||||
futures = [
|
||||
executor.submit(
|
||||
parse_context_length,
|
||||
data_file_lists[
|
||||
i * lines_for_each_th : (i + 1) * lines_for_each_th
|
||||
],
|
||||
data_type,
|
||||
i,
|
||||
)
|
||||
for i in range(task_num)
|
||||
]
|
||||
|
||||
for future in concurrent.futures.as_completed(futures):
|
||||
|
||||
json_dict[data_type].update(future.result())
|
||||
else:
|
||||
res = parse_context_length(data_file_lists, data_type)
|
||||
json_dict[data_type].update(res)
|
||||
|
||||
with open(jsonl_file_out, "w") as f:
|
||||
for key in json_dict[data_type_list[0]].keys():
|
||||
jsonl_line = {"key": key}
|
||||
for data_file in data_type_list:
|
||||
if key in json_dict[data_file]:
|
||||
jsonl_line.update(json_dict[data_file][key])
|
||||
jsonl_line = json.dumps(jsonl_line, ensure_ascii=False)
|
||||
f.write(jsonl_line + "\n")
|
||||
f.flush()
|
||||
print(f"processed {len(json_dict[data_type_list[0]])} samples")
|
||||
|
||||
else:
|
||||
pass
|
||||
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
|
||||
def parse_context_length(data_list: list, data_type: str, id=0):
|
||||
"""Parse context length.
|
||||
|
||||
Args:
|
||||
data_list: TODO.
|
||||
data_type: TODO.
|
||||
id: TODO.
|
||||
"""
|
||||
pbar = tqdm(total=len(data_list), dynamic_ncols=True)
|
||||
res = {}
|
||||
for i, line in enumerate(data_list):
|
||||
pbar.update(1)
|
||||
pbar.set_description(f"cpu: {id}")
|
||||
lines = line.strip().split(maxsplit=1)
|
||||
key = lines[0]
|
||||
line = lines[1] if len(lines) > 1 else ""
|
||||
line = line.strip()
|
||||
if data_type == "source":
|
||||
if os.path.exists(line):
|
||||
waveform, _ = librosa.load(line, sr=16000)
|
||||
sample_num = len(waveform)
|
||||
context_len = int(sample_num * 1000 / 16000 / 10)
|
||||
else:
|
||||
print("source file not found: {}".format(line))
|
||||
continue
|
||||
else:
|
||||
context_len = len(line.split()) if " " in line else len(line)
|
||||
res[key] = {data_type: line, f"{data_type}_len": context_len}
|
||||
return res
|
||||
|
||||
|
||||
@hydra.main(config_name=None, version_base=None)
|
||||
def main_hydra(cfg: DictConfig):
|
||||
|
||||
"""Main hydra.
|
||||
|
||||
Args:
|
||||
cfg: Configuration overrides.
|
||||
"""
|
||||
kwargs = OmegaConf.to_container(cfg, resolve=True)
|
||||
print(kwargs)
|
||||
|
||||
scp_file_list = kwargs.get(
|
||||
"scp_file_list",
|
||||
("/Users/zhifu/funasr1.0/test_local/wav.scp", "/Users/zhifu/funasr1.0/test_local/text.txt"),
|
||||
)
|
||||
if isinstance(scp_file_list, str):
|
||||
scp_file_list = eval(scp_file_list)
|
||||
data_type_list = kwargs.get("data_type_list", ("source", "target"))
|
||||
jsonl_file_out = kwargs.get(
|
||||
"jsonl_file_out", "/Users/zhifu/funasr1.0/test_local/audio_datasets.jsonl"
|
||||
)
|
||||
gen_jsonl_from_wav_text_list(
|
||||
scp_file_list, data_type_list=data_type_list, jsonl_file_out=jsonl_file_out
|
||||
)
|
||||
|
||||
|
||||
"""
|
||||
python -m funasr.datasets.audio_datasets.scp2jsonl \
|
||||
++scp_file_list='["/Users/zhifu/funasr1.0/test_local/wav.scp", "/Users/zhifu/funasr1.0/test_local/text.txt"]' \
|
||||
++data_type_list='["source", "target"]' \
|
||||
++jsonl_file_out=/Users/zhifu/funasr1.0/test_local/audio_datasets.jsonl
|
||||
"""
|
||||
|
||||
if __name__ == "__main__":
|
||||
main_hydra()
|
||||
@@ -0,0 +1,141 @@
|
||||
import os
|
||||
import json
|
||||
import torch
|
||||
import logging
|
||||
import hydra
|
||||
from omegaconf import DictConfig, OmegaConf
|
||||
import concurrent.futures
|
||||
import librosa
|
||||
import torch.distributed as dist
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
def gen_jsonl_from_wav_text_list(
|
||||
path, data_type_list=("source",), jsonl_file_out: str = None, **kwargs
|
||||
):
|
||||
"""Gen jsonl from wav text list.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
data_type_list: TODO.
|
||||
jsonl_file_out: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
try:
|
||||
rank = dist.get_rank()
|
||||
world_size = dist.get_world_size()
|
||||
except:
|
||||
rank = 0
|
||||
world_size = 1
|
||||
|
||||
cpu_cores = os.cpu_count() or 1
|
||||
print(f"convert wav.scp text to jsonl, ncpu: {cpu_cores}")
|
||||
if rank == 0:
|
||||
json_dict = {}
|
||||
# for data_type, data_file in zip(data_type_list, path):
|
||||
data_type = data_type_list[0]
|
||||
data_file = path
|
||||
json_dict[data_type] = {}
|
||||
with open(data_file, "r") as f:
|
||||
|
||||
data_file_lists = f.readlines()
|
||||
print("")
|
||||
lines_for_each_th = (len(data_file_lists) - 1) // cpu_cores + 1
|
||||
task_num = cpu_cores if len(data_file_lists) > cpu_cores else 1
|
||||
# import pdb;pdb.set_trace()
|
||||
if task_num > 1:
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=cpu_cores) as executor:
|
||||
|
||||
futures = [
|
||||
executor.submit(
|
||||
parse_context_length,
|
||||
data_file_lists[i * lines_for_each_th : (i + 1) * lines_for_each_th],
|
||||
data_type,
|
||||
i,
|
||||
)
|
||||
for i in range(task_num)
|
||||
]
|
||||
|
||||
for future in concurrent.futures.as_completed(futures):
|
||||
|
||||
json_dict[data_type].update(future.result())
|
||||
else:
|
||||
res = parse_context_length(data_file_lists, data_type)
|
||||
json_dict[data_type].update(res)
|
||||
|
||||
with open(jsonl_file_out, "w") as f:
|
||||
for key in json_dict[data_type_list[0]].keys():
|
||||
jsonl_line = {"key": key}
|
||||
for data_file in data_type_list:
|
||||
jsonl_line.update(json_dict[data_file][key])
|
||||
# jsonl_line = json.dumps(jsonl_line, ensure_ascii=False)
|
||||
source_len = jsonl_line["source_len"]
|
||||
jsonl_line = f"{key} {source_len}"
|
||||
f.write(jsonl_line + "\n")
|
||||
f.flush()
|
||||
print(f"processed {len(json_dict[data_type_list[0]])} samples")
|
||||
|
||||
else:
|
||||
pass
|
||||
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
|
||||
def parse_context_length(data_list: list, data_type: str, id=0):
|
||||
"""Parse context length.
|
||||
|
||||
Args:
|
||||
data_list: TODO.
|
||||
data_type: TODO.
|
||||
id: TODO.
|
||||
"""
|
||||
pbar = tqdm(total=len(data_list), dynamic_ncols=True)
|
||||
res = {}
|
||||
for i, line in enumerate(data_list):
|
||||
pbar.update(1)
|
||||
pbar.set_description(f"cpu: {id}")
|
||||
lines = line.strip().split(maxsplit=1)
|
||||
key = lines[0]
|
||||
line = lines[1] if len(lines) > 1 else ""
|
||||
line = line.strip()
|
||||
if os.path.exists(line):
|
||||
waveform, _ = librosa.load(line, sr=16000)
|
||||
sample_num = len(waveform)
|
||||
context_len = int(sample_num / 16000 * 1000 / 10)
|
||||
else:
|
||||
context_len = len(line.split()) if " " in line else len(line)
|
||||
res[key] = {data_type: line, f"{data_type}_len": context_len}
|
||||
return res
|
||||
|
||||
|
||||
@hydra.main(config_name=None, version_base=None)
|
||||
def main_hydra(cfg: DictConfig):
|
||||
|
||||
"""Main hydra.
|
||||
|
||||
Args:
|
||||
cfg: Configuration overrides.
|
||||
"""
|
||||
kwargs = OmegaConf.to_container(cfg, resolve=True)
|
||||
print(kwargs)
|
||||
|
||||
scp_file_list = kwargs.get("scp_file_list", "/Users/zhifu/funasr1.0/data/list/train_wav.scp")
|
||||
# if isinstance(scp_file_list, str):
|
||||
# scp_file_list = eval(scp_file_list)
|
||||
data_type_list = kwargs.get("data_type_list", ("source",))
|
||||
jsonl_file_out = kwargs.get("jsonl_file_out", "/Users/zhifu/funasr1.0/data/list/wav_len.txt")
|
||||
gen_jsonl_from_wav_text_list(
|
||||
scp_file_list, data_type_list=data_type_list, jsonl_file_out=jsonl_file_out
|
||||
)
|
||||
|
||||
|
||||
"""
|
||||
python -m funasr.datasets.audio_datasets.scp2jsonl \
|
||||
++scp_file_list='["/Users/zhifu/funasr1.0/test_local/wav.scp", "/Users/zhifu/funasr1.0/test_local/text.txt"]' \
|
||||
++data_type_list='["source", "target"]' \
|
||||
++jsonl_file_out=/Users/zhifu/funasr1.0/test_local/audio_datasets.jsonl
|
||||
"""
|
||||
|
||||
if __name__ == "__main__":
|
||||
main_hydra()
|
||||
@@ -0,0 +1,216 @@
|
||||
import os
|
||||
import json
|
||||
import torch
|
||||
import logging
|
||||
import hydra
|
||||
import re
|
||||
import string
|
||||
from omegaconf import DictConfig, OmegaConf
|
||||
import concurrent.futures
|
||||
import librosa
|
||||
import torch.distributed as dist
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
def gen_jsonl_from_wav_text_list(
|
||||
path, data_type_list=("source", "target"), jsonl_file_out: str = None, model_dir: str = "iic/SenseVoiceSmall", **kwargs
|
||||
):
|
||||
"""Gen jsonl from wav text list.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
data_type_list: TODO.
|
||||
jsonl_file_out: TODO.
|
||||
model_dir: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
try:
|
||||
rank = dist.get_rank()
|
||||
world_size = dist.get_world_size()
|
||||
except:
|
||||
rank = 0
|
||||
world_size = 1
|
||||
|
||||
cpu_cores = os.cpu_count() or 1
|
||||
print(f"convert wav.scp text to jsonl, ncpu: {cpu_cores}")
|
||||
if rank == 0:
|
||||
json_dict = {}
|
||||
for data_type, data_file in zip(data_type_list, path):
|
||||
json_dict[data_type] = {}
|
||||
with open(data_file, "r") as f:
|
||||
|
||||
data_file_lists = f.readlines()
|
||||
lines_for_each_th = (len(data_file_lists) - 1) // cpu_cores + 1
|
||||
task_num = cpu_cores if len(data_file_lists) > cpu_cores else 1
|
||||
# import pdb;pdb.set_trace()
|
||||
if task_num > 1:
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=cpu_cores) as executor:
|
||||
|
||||
futures = [
|
||||
executor.submit(
|
||||
parse_context_length,
|
||||
data_file_lists[
|
||||
i * lines_for_each_th : (i + 1) * lines_for_each_th
|
||||
],
|
||||
data_type,
|
||||
i,
|
||||
)
|
||||
for i in range(task_num)
|
||||
]
|
||||
|
||||
for future in concurrent.futures.as_completed(futures):
|
||||
|
||||
json_dict[data_type].update(future.result())
|
||||
else:
|
||||
res = parse_context_length(data_file_lists, data_type)
|
||||
json_dict[data_type].update(res)
|
||||
|
||||
if "text_language" not in data_type_list or "emo_target" not in data_type_list or "event_target" not in data_type_list:
|
||||
from funasr import AutoModel
|
||||
|
||||
model = AutoModel(
|
||||
model=model_dir,
|
||||
)
|
||||
|
||||
rich_dict = {}
|
||||
for key in json_dict["source"].keys():
|
||||
input_wav = json_dict["source"][key]["source"]
|
||||
res = model.generate(
|
||||
input=input_wav,
|
||||
cache={},
|
||||
language="auto", # "zn", "en", "yue", "ja", "ko", "nospeech"
|
||||
use_itn=True,
|
||||
)
|
||||
text = res[0]["text"]
|
||||
pattern = r"<\|[^|]+\|>"
|
||||
matches = re.findall(pattern, text)
|
||||
text_language, emo_target, event_target = matches[:3]
|
||||
rich_dict[key] = [text_language, emo_target, event_target]
|
||||
|
||||
|
||||
if "text_language" not in data_type_list:
|
||||
data_type_list.append("text_language")
|
||||
if "text_language" not in json_dict:
|
||||
json_dict["text_language"] = {}
|
||||
for key in json_dict["source"].keys():
|
||||
json_dict["text_language"][key] = {}
|
||||
json_dict["text_language"][key]["text_language"] = rich_dict[key][0]
|
||||
|
||||
if "emo_target" not in data_type_list:
|
||||
data_type_list.append("emo_target")
|
||||
if "emo_target" not in json_dict:
|
||||
json_dict["emo_target"] = {}
|
||||
for key in json_dict["source"].keys():
|
||||
json_dict["emo_target"][key] = {}
|
||||
json_dict["emo_target"][key]["emo_target"] = rich_dict[key][1]
|
||||
|
||||
if "event_target" not in data_type_list:
|
||||
data_type_list.append("event_target")
|
||||
if "event_target" not in json_dict:
|
||||
json_dict["event_target"] = {}
|
||||
for key in json_dict["source"].keys():
|
||||
json_dict["event_target"][key] = {}
|
||||
json_dict["event_target"][key]["event_target"] = rich_dict[key][2]
|
||||
|
||||
with open(jsonl_file_out, "w") as f:
|
||||
for key in json_dict[data_type_list[0]].keys():
|
||||
jsonl_line = {"key": key}
|
||||
for data_file in data_type_list:
|
||||
jsonl_line.update(json_dict[data_file][key])
|
||||
jsonl_line = json.dumps(jsonl_line, ensure_ascii=False)
|
||||
f.write(jsonl_line + "\n")
|
||||
f.flush()
|
||||
print(f"processed {len(json_dict[data_type_list[0]])} samples")
|
||||
|
||||
else:
|
||||
pass
|
||||
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
def contains_punctuation(s):
|
||||
"""Contains punctuation.
|
||||
|
||||
Args:
|
||||
s: TODO.
|
||||
"""
|
||||
punctuations = (
|
||||
string.punctuation +
|
||||
',。、;:?!""''()【】《》〈〉「」『』〔〕[]{}~·…—–'
|
||||
)
|
||||
return any(char in punctuations for char in s)
|
||||
|
||||
def parse_context_length(data_list: list, data_type: str, id=0):
|
||||
"""Parse context length.
|
||||
|
||||
Args:
|
||||
data_list: TODO.
|
||||
data_type: TODO.
|
||||
id: TODO.
|
||||
"""
|
||||
pbar = tqdm(total=len(data_list), dynamic_ncols=True)
|
||||
res = {}
|
||||
for i, line in enumerate(data_list):
|
||||
pbar.update(1)
|
||||
pbar.set_description(f"cpu: {id}")
|
||||
lines = line.strip().split(maxsplit=1)
|
||||
key = lines[0]
|
||||
line = lines[1] if len(lines) > 1 else ""
|
||||
line = line.strip()
|
||||
if os.path.exists(line):
|
||||
waveform, _ = librosa.load(line, sr=16000)
|
||||
sample_num = len(waveform)
|
||||
context_len = int(sample_num / 16000 * 1000 / 10)
|
||||
else:
|
||||
context_len = len(line.split()) if " " in line else len(line)
|
||||
if data_type == "source":
|
||||
res[key] = {data_type: line, f"{data_type}_len": context_len}
|
||||
elif data_type == "target":
|
||||
punc = contains_punctuation(line)
|
||||
if punc:
|
||||
with_or_wo_itn = "<|withitn|>"
|
||||
else:
|
||||
with_or_wo_itn = "<|woitn|>"
|
||||
res[key] = {data_type: line, f"{data_type}_len": context_len, "with_or_wo_itn": with_or_wo_itn}
|
||||
else:
|
||||
res[key] = {data_type: line}
|
||||
return res
|
||||
|
||||
|
||||
@hydra.main(config_name=None, version_base=None)
|
||||
def main_hydra(cfg: DictConfig):
|
||||
|
||||
"""Main hydra.
|
||||
|
||||
Args:
|
||||
cfg: Configuration overrides.
|
||||
"""
|
||||
kwargs = OmegaConf.to_container(cfg, resolve=True)
|
||||
print(kwargs)
|
||||
|
||||
scp_file_list = kwargs.get(
|
||||
"scp_file_list",
|
||||
("/Users/zhifu/funasr1.0/test_local/wav.scp", "/Users/zhifu/funasr1.0/test_local/text.txt"),
|
||||
)
|
||||
if isinstance(scp_file_list, str):
|
||||
scp_file_list = eval(scp_file_list)
|
||||
data_type_list = kwargs.get("data_type_list", ("source", "target"))
|
||||
jsonl_file_out = kwargs.get(
|
||||
"jsonl_file_out", "/Users/zhifu/funasr1.0/test_local/audio_datasets.jsonl"
|
||||
)
|
||||
model_dir = kwargs.get("model_dir", "iic/SenseVoiceSmall")
|
||||
gen_jsonl_from_wav_text_list(
|
||||
scp_file_list, data_type_list=data_type_list, jsonl_file_out=jsonl_file_out, model_dir=model_dir
|
||||
)
|
||||
|
||||
|
||||
"""
|
||||
python -m funasr.datasets.audio_datasets.sensevoice2jsonl \
|
||||
++scp_file_list='["/Users/zhifu/funasr1.0/test_local/wav.scp", "/Users/zhifu/funasr1.0/test_local/text.txt", "/Users/zhifu/funasr1.0/test_local/text_language.txt", "/Users/zhifu/funasr1.0/test_local/emo_target.txt", "/Users/zhifu/funasr1.0/test_local/event_target.txt"]' \
|
||||
++data_type_list='["source", "target", "text_language", "emo_target", "event_target"]' \
|
||||
++jsonl_file_out='/Users/zhifu/funasr1.0/test_local/audio_datasets.jsonl' \
|
||||
++model_dir='iic/SenseVoiceSmall'
|
||||
"""
|
||||
|
||||
if __name__ == "__main__":
|
||||
main_hydra()
|
||||
@@ -0,0 +1,124 @@
|
||||
import os
|
||||
import json
|
||||
import torch
|
||||
import logging
|
||||
import hydra
|
||||
from omegaconf import DictConfig, OmegaConf
|
||||
import concurrent.futures
|
||||
import librosa
|
||||
import torch.distributed as dist
|
||||
import threading
|
||||
from tqdm import tqdm
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
|
||||
def gen_scp_from_jsonl(jsonl_file, jsonl_file_out, ncpu):
|
||||
"""Gen scp from jsonl.
|
||||
|
||||
Args:
|
||||
jsonl_file: TODO.
|
||||
jsonl_file_out: TODO.
|
||||
ncpu: TODO.
|
||||
"""
|
||||
jsonl_file_out_f = open(jsonl_file_out, "w")
|
||||
with open(jsonl_file, encoding="utf-8") as fin:
|
||||
lines = fin.readlines()
|
||||
|
||||
num_total = len(lines)
|
||||
if ncpu > 1:
|
||||
# 使用ThreadPoolExecutor限制并发线程数
|
||||
with ThreadPoolExecutor(max_workers=ncpu) as executor:
|
||||
# 提交任务到线程池
|
||||
futures = {executor.submit(update_data, lines, i) for i in tqdm(range(num_total))}
|
||||
|
||||
# 等待所有任务完成,这会阻塞直到所有提交的任务完成
|
||||
for future in concurrent.futures.as_completed(futures):
|
||||
# 这里可以添加额外的逻辑来处理完成的任务,但在这个例子中我们只是等待
|
||||
pass
|
||||
else:
|
||||
for i in range(num_total):
|
||||
update_data(lines, i)
|
||||
logging.info("All audio durations have been processed.")
|
||||
|
||||
for line in lines:
|
||||
|
||||
jsonl_file_out_f.write(line + "\n")
|
||||
jsonl_file_out_f.flush()
|
||||
|
||||
jsonl_file_out_f.close()
|
||||
|
||||
|
||||
def update_data(lines, i):
|
||||
"""Update data.
|
||||
|
||||
Args:
|
||||
lines: TODO.
|
||||
i: TODO.
|
||||
"""
|
||||
line = lines[i]
|
||||
data = json.loads(line.strip())
|
||||
|
||||
wav_path = data["source"].replace("/cpfs01", "/cpfs_speech/data")
|
||||
if os.path.exists(wav_path):
|
||||
waveform, _ = librosa.load(wav_path, sr=16000)
|
||||
sample_num = len(waveform)
|
||||
source_len = int(sample_num / 16000 * 1000 / 10)
|
||||
source_len_old = data["source_len"]
|
||||
# if (source_len_old - source_len) > 100 or (source_len - source_len_old) > 100:
|
||||
# logging.info(f"old: {source_len_old}, new: {source_len}, wav: {wav_path}")
|
||||
data["source_len"] = source_len
|
||||
data["source"] = wav_path
|
||||
jsonl_line = json.dumps(data, ensure_ascii=False)
|
||||
lines[i] = jsonl_line
|
||||
|
||||
|
||||
def update_wav_len(jsonl_file_list_in, jsonl_file_out_dir, ncpu=1):
|
||||
|
||||
"""Update wav len.
|
||||
|
||||
Args:
|
||||
jsonl_file_list_in: TODO.
|
||||
jsonl_file_out_dir: TODO.
|
||||
ncpu: TODO.
|
||||
"""
|
||||
os.makedirs(jsonl_file_out_dir, exist_ok=True)
|
||||
with open(jsonl_file_list_in, "r") as f:
|
||||
data_file_lists = f.readlines()
|
||||
|
||||
for i, jsonl in enumerate(data_file_lists):
|
||||
filename_with_extension = os.path.basename(jsonl.strip())
|
||||
jsonl_file_out = os.path.join(jsonl_file_out_dir, filename_with_extension)
|
||||
logging.info(f"{i}/{len(data_file_lists)}, jsonl: {jsonl}, {jsonl_file_out}")
|
||||
|
||||
gen_scp_from_jsonl(jsonl.strip(), jsonl_file_out, ncpu)
|
||||
|
||||
|
||||
@hydra.main(config_name=None, version_base=None)
|
||||
def main_hydra(cfg: DictConfig):
|
||||
|
||||
"""Main hydra.
|
||||
|
||||
Args:
|
||||
cfg: Configuration overrides.
|
||||
"""
|
||||
kwargs = OmegaConf.to_container(cfg, resolve=True)
|
||||
logging.info(kwargs)
|
||||
|
||||
jsonl_file_list_in = kwargs.get(
|
||||
"jsonl_file_list_in", "/Users/zhifu/funasr1.0/data/list/data_jsonl.list"
|
||||
)
|
||||
jsonl_file_out_dir = kwargs.get("jsonl_file_out_dir", "/Users/zhifu/funasr1.0/data_tmp")
|
||||
ncpu = kwargs.get("ncpu", 1)
|
||||
update_wav_len(jsonl_file_list_in, jsonl_file_out_dir, ncpu)
|
||||
# gen_scp_from_jsonl(jsonl_file_list_in, jsonl_file_out_dir)
|
||||
|
||||
|
||||
"""
|
||||
python -m funasr.datasets.audio_datasets.json2scp \
|
||||
++scp_file_list='["/Users/zhifu/funasr1.0/test_local/wav.scp", "/Users/zhifu/funasr1.0/test_local/text.txt"]' \
|
||||
++data_type_list='["source", "target"]' \
|
||||
++jsonl_file_in=/Users/zhifu/funasr1.0/test_local/audio_datasets.jsonl
|
||||
"""
|
||||
|
||||
if __name__ == "__main__":
|
||||
main_hydra()
|
||||
@@ -0,0 +1,168 @@
|
||||
import logging
|
||||
import torch
|
||||
|
||||
from funasr.register import tables
|
||||
|
||||
|
||||
# @tables.register("dataloader_classes", "DataloaderMapStyle")
|
||||
def DataloaderMapStyle(frontend=None, tokenizer=None, **kwargs):
|
||||
# dataset
|
||||
"""Dataloadermapstyle.
|
||||
|
||||
Args:
|
||||
frontend: Audio frontend for feature extraction.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
logging.info("Build dataloader")
|
||||
dataset_class = tables.dataset_classes.get(kwargs.get("dataset", "AudioDataset"))
|
||||
dataset_tr = dataset_class(
|
||||
kwargs.get("train_data_set_list"),
|
||||
frontend=frontend,
|
||||
tokenizer=tokenizer,
|
||||
is_training=True,
|
||||
**kwargs.get("dataset_conf"),
|
||||
)
|
||||
dataset_val = dataset_class(
|
||||
kwargs.get("valid_data_set_list"),
|
||||
frontend=frontend,
|
||||
tokenizer=tokenizer,
|
||||
is_training=False,
|
||||
**kwargs.get("dataset_conf"),
|
||||
)
|
||||
|
||||
# dataloader
|
||||
batch_sampler = kwargs["dataset_conf"].get("batch_sampler", "BatchSampler")
|
||||
batch_sampler_val = None
|
||||
if batch_sampler is not None:
|
||||
batch_sampler_class = tables.batch_sampler_classes.get(batch_sampler)
|
||||
batch_sampler = batch_sampler_class(dataset_tr, **kwargs.get("dataset_conf"))
|
||||
batch_sampler_val = batch_sampler_class(
|
||||
dataset_val, is_training=False, **kwargs.get("dataset_conf")
|
||||
)
|
||||
|
||||
dataloader_tr = torch.utils.data.DataLoader(
|
||||
dataset_tr, collate_fn=dataset_tr.collator, **batch_sampler
|
||||
)
|
||||
dataloader_val = torch.utils.data.DataLoader(
|
||||
dataset_val, collate_fn=dataset_val.collator, **batch_sampler_val
|
||||
)
|
||||
|
||||
return dataloader_tr, dataloader_val
|
||||
|
||||
|
||||
@tables.register("dataloader_classes", "DataloaderMapStyle")
|
||||
class DataloaderMapStyle:
|
||||
def __init__(self, frontend=None, tokenizer=None, **kwargs):
|
||||
# dataset
|
||||
"""Initialize DataloaderMapStyle.
|
||||
|
||||
Args:
|
||||
frontend: Audio frontend for feature extraction.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
logging.info("Build dataloader")
|
||||
|
||||
dataset_class = tables.dataset_classes.get(kwargs.get("dataset", "AudioDataset"))
|
||||
dataset_tr = None
|
||||
# split dataset
|
||||
self.data_split_num = kwargs["dataset_conf"].get("data_split_num", 1)
|
||||
if self.data_split_num == 1:
|
||||
dataset_tr = dataset_class(
|
||||
kwargs.get("train_data_set_list"),
|
||||
frontend=frontend,
|
||||
tokenizer=tokenizer,
|
||||
is_training=True,
|
||||
**kwargs.get("dataset_conf"),
|
||||
)
|
||||
dataset_val = dataset_class(
|
||||
kwargs.get("valid_data_set_list"),
|
||||
frontend=frontend,
|
||||
tokenizer=tokenizer,
|
||||
is_training=False,
|
||||
**kwargs.get("dataset_conf"),
|
||||
)
|
||||
|
||||
self.dataset_tr = dataset_tr
|
||||
self.dataset_val = dataset_val
|
||||
self.kwargs = kwargs
|
||||
|
||||
self.dataset_class = dataset_class
|
||||
self.frontend = frontend
|
||||
self.tokenizer = tokenizer
|
||||
self.kwargs = kwargs
|
||||
|
||||
def build_iter(self, epoch=0, data_split_i=0, start_step=0, **kwargs):
|
||||
|
||||
# reload dataset slice
|
||||
"""Build iter.
|
||||
|
||||
Args:
|
||||
epoch: TODO.
|
||||
data_split_i: TODO.
|
||||
start_step: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
if self.data_split_num > 1:
|
||||
del self.dataset_tr
|
||||
self.dataset_tr = self.dataset_class(
|
||||
self.kwargs.get("train_data_set_list"),
|
||||
frontend=self.frontend,
|
||||
tokenizer=self.tokenizer,
|
||||
is_training=True,
|
||||
**self.kwargs.get("dataset_conf"),
|
||||
data_split_i=data_split_i,
|
||||
)
|
||||
|
||||
# dataloader
|
||||
batch_sampler = self.kwargs["dataset_conf"].get("batch_sampler", "BatchSampler")
|
||||
batch_sampler_val = None
|
||||
if batch_sampler is not None:
|
||||
batch_sampler_class = tables.batch_sampler_classes.get(batch_sampler)
|
||||
batch_sampler = batch_sampler_class(
|
||||
self.dataset_tr, start_step=start_step, **self.kwargs.get("dataset_conf")
|
||||
)
|
||||
batch_sampler_val = batch_sampler_class(
|
||||
self.dataset_val, is_training=False, **self.kwargs.get("dataset_conf")
|
||||
)
|
||||
|
||||
batch_sampler["batch_sampler"].set_epoch(epoch)
|
||||
batch_sampler_val["batch_sampler"].set_epoch(epoch)
|
||||
dataloader_tr = torch.utils.data.DataLoader(
|
||||
self.dataset_tr, collate_fn=self.dataset_tr.collator, **batch_sampler
|
||||
)
|
||||
dataloader_val = torch.utils.data.DataLoader(
|
||||
self.dataset_val, collate_fn=self.dataset_val.collator, **batch_sampler_val
|
||||
)
|
||||
|
||||
return dataloader_tr, dataloader_val
|
||||
|
||||
|
||||
@tables.register("dataloader_classes", "DataloaderIterable")
|
||||
def DataloaderIterable(frontend=None, tokenizer=None, **kwargs):
|
||||
"""Dataloaderiterable.
|
||||
|
||||
Args:
|
||||
frontend: Audio frontend for feature extraction.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
logging.info("Build dataloader")
|
||||
dataset_class = tables.dataset_classes.get(kwargs.get("dataset", "LargeDataset"))
|
||||
dataset_tr = dataset_class(
|
||||
kwargs.get("train_data_set_list"),
|
||||
frontend=frontend,
|
||||
tokenizer=tokenizer,
|
||||
is_training=True,
|
||||
**kwargs.get("dataset_conf"),
|
||||
)
|
||||
dataset_val = dataset_class(
|
||||
kwargs.get("valid_data_set_list"),
|
||||
frontend=frontend,
|
||||
tokenizer=tokenizer,
|
||||
is_training=False,
|
||||
**kwargs.get("dataset_conf"),
|
||||
)
|
||||
|
||||
return dataset_tr, dataset_val
|
||||
@@ -0,0 +1,569 @@
|
||||
import json
|
||||
import logging
|
||||
|
||||
import re
|
||||
import torch
|
||||
import random
|
||||
import traceback
|
||||
import numpy as np
|
||||
from funasr.register import tables
|
||||
from funasr.utils.load_utils import extract_fbank, load_audio_text_image_video
|
||||
|
||||
|
||||
@tables.register("dataset_classes", "FunASR")
|
||||
class FunASR(torch.utils.data.Dataset):
|
||||
"""
|
||||
FunASR dataset
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path,
|
||||
index_ds: str = None,
|
||||
frontend=None,
|
||||
tokenizer=None,
|
||||
int_pad_value: int = -1,
|
||||
float_pad_value: float = 0.0,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize FunASR.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
index_ds: TODO.
|
||||
frontend: Audio frontend for feature extraction.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
int_pad_value: TODO.
|
||||
float_pad_value: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
index_ds_class = tables.index_ds_classes.get(index_ds)
|
||||
self.index_ds = index_ds_class(path, **kwargs)
|
||||
preprocessor_speech = kwargs.get("preprocessor_speech", None)
|
||||
if preprocessor_speech:
|
||||
preprocessor_speech_class = tables.preprocessor_classes.get(preprocessor_speech)
|
||||
preprocessor_speech = preprocessor_speech_class(
|
||||
**kwargs.get("preprocessor_speech_conf")
|
||||
)
|
||||
self.preprocessor_speech = preprocessor_speech
|
||||
|
||||
preprocessor_noise = kwargs.get("preprocessor_noise", None)
|
||||
if preprocessor_noise:
|
||||
preprocessor_noise_class = tables.preprocessor_classes.get(preprocessor_noise)
|
||||
preprocessor_noise = preprocessor_noise_class(**kwargs.get("preprocessor_noise_conf"))
|
||||
self.preprocessor_noise = preprocessor_noise
|
||||
|
||||
prompt_classes_text = kwargs.get("prompt_classes", None)
|
||||
if prompt_classes_text is not None:
|
||||
prompt_classes = tables.prompt_classes.get(prompt_classes_text)
|
||||
prompt_classes = prompt_classes(**kwargs.get("prompt_conf"))
|
||||
else:
|
||||
prompt_classes = None
|
||||
self.prompt_classes = prompt_classes
|
||||
|
||||
self.frontend = frontend
|
||||
self.fs = 16000 if frontend is None else frontend.fs
|
||||
self.data_type = "sound"
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.int_pad_value = int_pad_value
|
||||
self.float_pad_value = float_pad_value
|
||||
self.sos = kwargs.get("sos", "<|startoftranscript|>")
|
||||
self.eos = kwargs.get("eos", "<|endoftext|>")
|
||||
self.batch_size = kwargs.get("batch_size")
|
||||
self.batch_type = kwargs.get("batch_type")
|
||||
self.prompt_ids_len = 0
|
||||
self.retry = kwargs.get("retry", 100)
|
||||
|
||||
self.pattern = re.compile(r"(<\|startofspeech\|>.*?<\|endofspeech\|>)")
|
||||
# self.kwargs = kwargs
|
||||
self.max_token_length = kwargs.get("max_token_length", 1500)
|
||||
self.batch_size_scale_ratio_max = kwargs.get("batch_size_scale_ratio_max", 1.5)
|
||||
self.batch_size_token_max = kwargs.get("batch_size_token_max", 2500)
|
||||
self.multiturn_num_max = kwargs.get("multiturn_num_max", 5)
|
||||
self.max_source_length = kwargs.get("max_source_length", 3000)
|
||||
self.max_target_length = kwargs.get("max_target_length", 1024)
|
||||
self.do_think = kwargs.get("do_think", True)
|
||||
self.sys_prompt = kwargs.get("sys_prompt", True)
|
||||
|
||||
# used for dynamic output alignment
|
||||
self.use_dynamic_output_ratio = kwargs.get("use_dynamic_output_ratio", 0.0)
|
||||
self.min_output_mask_token_len = kwargs.get("min_mask_token_len", 1)
|
||||
self.min_output_non_mask_token_len = kwargs.get("min_non_mask_token_len", 6) # [eos]
|
||||
|
||||
def get_source_len(self, index):
|
||||
"""Get source len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_source_len(item)
|
||||
|
||||
def get_target_len(self, index):
|
||||
"""Get target len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_target_len(item)
|
||||
|
||||
def get_random_user_prompt(self, item, user_prompt):
|
||||
"""Get random user prompt.
|
||||
|
||||
Args:
|
||||
item: TODO.
|
||||
user_prompt: TODO.
|
||||
"""
|
||||
tasks = ["语音转写:", "Speech transcription:"]
|
||||
language = item.get("language", None)
|
||||
# LID in distill data is fake
|
||||
language = None
|
||||
if language is not None:
|
||||
if language.lower() == "zh":
|
||||
tasks.append("语音转写成中文:")
|
||||
tasks.append("Transcribe speech into Chinese:")
|
||||
elif language.lower() == "en":
|
||||
tasks.append("语音转写成英文:")
|
||||
tasks.append("Transcribe speech into English:")
|
||||
if len(tasks) == 2:
|
||||
task = random.choice(tasks)
|
||||
elif len(tasks) == 4:
|
||||
task = random.choices(tasks, weights=[0.4, 0.4, 0.1, 0.1])[0]
|
||||
if "语音转写:<|startofspeech|>" in user_prompt:
|
||||
user_prompt = user_prompt.replace("语音转写:<|startofspeech|>", task + "<|startofspeech|>")
|
||||
elif "Speech transcription:<|startofspeech|>" in user_prompt:
|
||||
user_prompt = user_prompt.replace("Speech transcription:<|startofspeech|>", task + "<|startofspeech|>")
|
||||
return user_prompt
|
||||
|
||||
def __len__(self):
|
||||
"""Internal: len ."""
|
||||
return len(self.index_ds)
|
||||
|
||||
def __getitem__(self, index):
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
output = None
|
||||
|
||||
for idx in range(self.retry):
|
||||
if idx > 0:
|
||||
logging.info(f"retry: {idx}")
|
||||
badcase_flag = False
|
||||
if idx == 0:
|
||||
index_cur = index
|
||||
else:
|
||||
index_cur = torch.randint(0, len(self.index_ds), ()).item()
|
||||
|
||||
item = self.index_ds[index_cur]
|
||||
|
||||
system = item["system"]
|
||||
user = item["user"]
|
||||
assistant = item["assistant"]
|
||||
is_noised = item.get("noised", False)
|
||||
if len(user) < 1 or len(assistant) < 1:
|
||||
logging.warning(f"item is error: {item}")
|
||||
continue
|
||||
input_ids, labels, fbank, fbank_lens, fbank_mask, fbank_beg, fake_token_len = (
|
||||
[],
|
||||
[],
|
||||
[],
|
||||
[],
|
||||
[],
|
||||
[],
|
||||
[],
|
||||
)
|
||||
|
||||
for i, (system_prompt, user_prompt, target_out) in enumerate(
|
||||
zip(system, user, assistant)
|
||||
):
|
||||
if i >= self.multiturn_num_max:
|
||||
break
|
||||
if len(input_ids) > self.max_token_length:
|
||||
logging.info(
|
||||
f"input_ids > max_token_length: {len(input_ids)}>{self.max_token_length}, {item}"
|
||||
)
|
||||
break
|
||||
|
||||
if self.prompt_classes is not None:
|
||||
asr_prompt = user_prompt.split("<|startofspeech|>")[0]
|
||||
language = self.prompt_classes.detect_language(asr_prompt)
|
||||
user_prompt_all_context = self.prompt_classes.get_prompt(item, language)
|
||||
else:
|
||||
user_prompt_all_context = ""
|
||||
|
||||
if i == 0:
|
||||
source_input = f"<|im_start|>system\n{system_prompt}<|im_end|>\n<|im_start|>user\n{user_prompt_all_context}{user_prompt}<|im_end|>\n<|im_start|>assistant\n"
|
||||
if not self.sys_prompt:
|
||||
source_input = f"<|im_start|>user\n{user_prompt_all_context}{user_prompt}<|im_end|>\n<|im_start|>assistant\n"
|
||||
else:
|
||||
source_input = (
|
||||
f"<|im_start|>user\n{user_prompt}<|im_end|>\n<|im_start|>assistant\n"
|
||||
)
|
||||
if not self.do_think:
|
||||
source_input += "<think>\n\n</think>\n\n"
|
||||
splits = self.pattern.split(source_input)
|
||||
source_ids = []
|
||||
fbank_i = []
|
||||
fake_token_len_i = 0
|
||||
fbank_beg_i = -1
|
||||
fbank_lens_i = []
|
||||
speech = []
|
||||
speech_lengths = []
|
||||
for k, sub_str in enumerate(splits):
|
||||
if not sub_str.startswith("<|startofspeech|>"):
|
||||
sub_token = self.tokenizer.encode(sub_str)
|
||||
source_ids += sub_token
|
||||
else:
|
||||
sub_str = sub_str.replace("<|startofspeech|>", "").replace(
|
||||
"<|endofspeech|>", ""
|
||||
)
|
||||
if sub_str.startswith("!"):
|
||||
try:
|
||||
data_src = load_audio_text_image_video(sub_str[1:], fs=self.fs)
|
||||
if self.preprocessor_noise is not None and not is_noised:
|
||||
try:
|
||||
data_src = self.preprocessor_noise(data_src.numpy())
|
||||
except Exception as e:
|
||||
logging.error(f"Generate noise audio failed: {e}")
|
||||
|
||||
speech, speech_lengths = extract_fbank(
|
||||
data_src,
|
||||
data_type=self.data_type,
|
||||
frontend=self.frontend,
|
||||
is_final=True,
|
||||
) # speech: [b, T, d]
|
||||
if speech_lengths > self.max_source_length:
|
||||
logging.info(
|
||||
f"speech_lengths > max_source_length: {speech_lengths}>{self.max_source_length}, {item}"
|
||||
)
|
||||
badcase_flag = True
|
||||
except Exception as e:
|
||||
logging.warning(
|
||||
f"Loading wav failed! {str(e)}, {traceback.format_exc()}\n{item}"
|
||||
)
|
||||
badcase_flag = True
|
||||
continue
|
||||
if True:
|
||||
olens = 1 + (speech_lengths[0].item() - 3 + 2 * 1) // 2
|
||||
olens = 1 + (olens - 3 + 2 * 1) // 2
|
||||
fake_token_len_i = (olens - 1) // 2 + 1
|
||||
else:
|
||||
fake_token_len_i = speech_lengths[0].item()
|
||||
fake_token = [0] * fake_token_len_i
|
||||
fbank_beg_i = len(source_ids)
|
||||
source_ids += fake_token
|
||||
|
||||
if badcase_flag:
|
||||
continue
|
||||
if fbank_beg_i > 0:
|
||||
fbank_beg += [fbank_beg_i + len(input_ids)]
|
||||
fake_token_len += [fake_token_len_i]
|
||||
else:
|
||||
fbank_beg += [-1]
|
||||
fake_token_len += [0]
|
||||
|
||||
if target_out is not None and any(
|
||||
isinstance(item, dict) and "prev_content" in item for item in target_out
|
||||
):
|
||||
prev_value = next(
|
||||
(
|
||||
item["prev_content"]
|
||||
for item in target_out
|
||||
if isinstance(item, dict) and "prev_content" in item
|
||||
),
|
||||
None,
|
||||
)
|
||||
source_ids += self.tokenizer.encode(prev_value)
|
||||
source_mask = [-100] * len(source_ids)
|
||||
target_out = f"{target_out[0]}<|im_end|>"
|
||||
else:
|
||||
source_mask = [-100] * len(source_ids)
|
||||
target_out = f"{target_out}<|im_end|>"
|
||||
target_ids = self.tokenizer.encode(target_out)
|
||||
|
||||
if len(target_ids) > self.max_target_length:
|
||||
logging.info(
|
||||
f"text_length: {len(target_ids)} > {self.max_target_length}, drop it: {item}"
|
||||
)
|
||||
# simulate prev-token fixed output
|
||||
target_labels = target_ids.copy()
|
||||
if np.random.rand() < self.use_dynamic_output_ratio:
|
||||
max_len = len(target_labels)
|
||||
min_output_mask_token_len = min(self.min_output_mask_token_len, max_len)
|
||||
min_output_non_mask_token_len = min(self.min_output_non_mask_token_len, max_len)
|
||||
if max_len - min_output_non_mask_token_len > min_output_mask_token_len:
|
||||
end_index = np.random.randint(min_output_mask_token_len,
|
||||
max_len - min_output_non_mask_token_len)
|
||||
else:
|
||||
end_index = max_len - min_output_non_mask_token_len
|
||||
if end_index > 0:
|
||||
target_labels[:end_index] = [-100] * end_index
|
||||
|
||||
input_ids += source_ids + target_ids
|
||||
labels += source_mask + target_labels
|
||||
if len(speech) > 0:
|
||||
fbank.append(speech[0, :, :])
|
||||
fbank_lens.append(speech_lengths)
|
||||
if badcase_flag:
|
||||
continue
|
||||
|
||||
input_ids = torch.tensor(input_ids, dtype=torch.int64) # [: self.max_token_length]
|
||||
attention_mask = torch.tensor([1] * len(input_ids), dtype=torch.int32)
|
||||
labels = torch.tensor(labels, dtype=torch.int64) # [: self.max_token_length]
|
||||
|
||||
fbank_beg = torch.tensor(fbank_beg, dtype=torch.int32)
|
||||
fake_token_len = torch.tensor(fake_token_len, dtype=torch.int32)
|
||||
|
||||
output = {
|
||||
"fbank_beg": fbank_beg,
|
||||
"fake_token_len": fake_token_len,
|
||||
"input_ids": input_ids,
|
||||
"attention_mask": attention_mask,
|
||||
"labels_ids": labels,
|
||||
}
|
||||
output["item"] = item
|
||||
if len(fbank) > 0:
|
||||
output["speech"] = fbank
|
||||
output["speech_lengths"] = fbank_lens
|
||||
if len(input_ids) > self.max_token_length:
|
||||
logging.warning(
|
||||
f"len(input_ids): {len(input_ids)} > max_token_length: {self.max_token_length}, item: {item}"
|
||||
)
|
||||
continue
|
||||
|
||||
break
|
||||
|
||||
return output
|
||||
|
||||
def collator(self, samples: list = None):
|
||||
|
||||
"""Collator.
|
||||
|
||||
Args:
|
||||
samples: TODO.
|
||||
"""
|
||||
for idx in range(self.retry):
|
||||
badcase_flag = False
|
||||
|
||||
outputs = {}
|
||||
for sample in samples:
|
||||
if sample is None:
|
||||
continue
|
||||
for key in sample.keys():
|
||||
if key not in outputs:
|
||||
outputs[key] = []
|
||||
if isinstance(sample[key], (list, tuple)):
|
||||
outputs[key].extend(sample[key])
|
||||
else:
|
||||
outputs[key].append(sample[key])
|
||||
|
||||
for key, data_list in outputs.items():
|
||||
if isinstance(data_list[0], torch.Tensor):
|
||||
if data_list[0].dtype == torch.int64 or data_list[0].dtype == torch.int32:
|
||||
|
||||
pad_value = self.int_pad_value
|
||||
else:
|
||||
pad_value = self.float_pad_value
|
||||
|
||||
outputs[key] = torch.nn.utils.rnn.pad_sequence(
|
||||
data_list, batch_first=True, padding_value=pad_value
|
||||
)
|
||||
|
||||
if self.batch_type != "example":
|
||||
b, t = outputs["input_ids"].shape
|
||||
if b > 1 and b * t > self.batch_size_token_max:
|
||||
logging.info(
|
||||
f"Warning, {idx}th, b*t: {b}*{t}={b * t} > batch_size_sample_max: {self.batch_size_token_max}, drop last data"
|
||||
)
|
||||
samples = samples[:-1]
|
||||
continue
|
||||
|
||||
break
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
@tables.register("index_ds_classes", "FunASR")
|
||||
class FunASR(torch.utils.data.Dataset): # torch.utils.data.Dataset
|
||||
|
||||
def __init__(self, path: str, **kwargs):
|
||||
"""Initialize FunASR.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
self.max_source_length = kwargs.get("max_source_length", 8000)
|
||||
self.min_source_length = kwargs.get("min_source_length", 10)
|
||||
self.max_target_length = kwargs.get("max_target_length", 2048)
|
||||
self.min_target_length = kwargs.get("min_target_length", 0)
|
||||
# self.max_token_length = kwargs.get("max_token_length", 2200)+
|
||||
audio_downsample_rate = int(kwargs.get("audio_downsample_rate", 8))
|
||||
|
||||
is_training = kwargs.get("is_training", True)
|
||||
if not (path.endswith(".jsonl") or path.endswith(".json")):
|
||||
# jsonl list file
|
||||
data_split_num = kwargs.get("data_split_num", 1)
|
||||
data_split_i = kwargs.get("data_split_i", 0)
|
||||
|
||||
if not is_training:
|
||||
data_split_num = 1
|
||||
data_split_i = 0
|
||||
with open(path, encoding="utf-8") as fin:
|
||||
file_list_all = fin.readlines()
|
||||
|
||||
num_per_slice = (len(file_list_all) - 1) // data_split_num + 1 # 16
|
||||
file_list = file_list_all[
|
||||
data_split_i * num_per_slice: (data_split_i + 1) * num_per_slice
|
||||
]
|
||||
logging.info(
|
||||
f"is_training: {is_training}, data_split_num: {data_split_num}, data_split_i: {data_split_i}, \nfile_list: {file_list}, \nfile_list_all: {file_list_all}"
|
||||
)
|
||||
|
||||
else:
|
||||
file_list = [path]
|
||||
|
||||
contents = []
|
||||
total_whrs = 0.0
|
||||
total_token_for_llm_B = 0.0
|
||||
for file_json in file_list:
|
||||
with open(file_json.strip(), encoding="utf-8") as fin:
|
||||
for line in fin:
|
||||
try:
|
||||
data_dict = json.loads(line.strip())
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
f"drop it, json error: {e}, line: {line}, file_json: {file_json}"
|
||||
)
|
||||
continue
|
||||
|
||||
data = data_dict["messages"]
|
||||
if isinstance(data_dict.get("speech_length", 0), (list, tuple)):
|
||||
speech_length = int(data_dict.get("speech_length", [0])[0])
|
||||
text_length = int(data_dict.get("text_length", 0)[0])
|
||||
else:
|
||||
speech_length = int(data_dict.get("speech_length", 0))
|
||||
text_length = int(data_dict.get("text_length", 0))
|
||||
speech_length = int(speech_length)
|
||||
text_length = int(text_length)
|
||||
if speech_length > 0 and speech_length < 1:
|
||||
continue
|
||||
if text_length < 1:
|
||||
logging.warning(
|
||||
f"speech_length: {speech_length}, text_length: {text_length}, data: {data}, file_json: {file_json}"
|
||||
)
|
||||
if len(data) > 2:
|
||||
text_length = len(data[2]['content'])
|
||||
continue
|
||||
if speech_length > self.max_source_length:
|
||||
continue
|
||||
if speech_length < self.min_source_length:
|
||||
continue
|
||||
if text_length > self.max_target_length:
|
||||
continue
|
||||
|
||||
system, user, assistant = [], [], []
|
||||
for i, item in enumerate(data):
|
||||
try:
|
||||
role = item["role"]
|
||||
content = item["content"]
|
||||
except KeyError:
|
||||
logging.error(
|
||||
f"drop it, KeyError: {item}, file_json: {file_json}"
|
||||
)
|
||||
continue
|
||||
|
||||
if role == "system":
|
||||
system.append(content)
|
||||
elif role == "user":
|
||||
user.append(content)
|
||||
elif role == "assistant":
|
||||
if "prev_content" in item:
|
||||
prev_content = item["prev_content"]
|
||||
assistant.append([content, {"prev_content": prev_content}])
|
||||
else:
|
||||
assistant.append(content)
|
||||
if len(system) == 0:
|
||||
system = ["You are a helpful assistant."]
|
||||
system = system * len(user)
|
||||
|
||||
contents_i = {
|
||||
"system": system,
|
||||
"user": user,
|
||||
"assistant": assistant,
|
||||
"source_len": speech_length + text_length,
|
||||
}
|
||||
if "key" in data_dict:
|
||||
contents_i["key"] = data_dict["key"] if not isinstance(data_dict.get("key", "key_01234"),
|
||||
(list, tuple)) else data_dict["key"][0]
|
||||
|
||||
if "hist_context" in data_dict:
|
||||
contents_i["hist_context"] = data_dict["hist_context"]
|
||||
if "hotwords" in data_dict:
|
||||
contents_i["hotwords"] = data_dict["hotwords"]
|
||||
if "asr_hotwords" in data_dict:
|
||||
contents_i["asr_hotwords"] = data_dict["asr_hotwords"]
|
||||
if "vad_segs" in data_dict:
|
||||
contents_i["vad_segs"] = data_dict["vad_segs"]
|
||||
if "word_list" in data_dict:
|
||||
contents_i["word_list"] = data_dict["word_list"]
|
||||
if "one_pass_result" in data_dict:
|
||||
contents_i["one_pass_result"] = data_dict["one_pass_result"]
|
||||
if "one_pass_wer" in data_dict:
|
||||
contents_i["one_pass_wer"] = data_dict["one_pass_wer"]
|
||||
if "noised" in data_dict:
|
||||
contents_i["noised"] = data_dict["noised"]
|
||||
|
||||
if kwargs.get("save_meta", False):
|
||||
contents_i["meta"] = data_dict
|
||||
|
||||
total_whrs += speech_length / 100.0 / 3600 / 10000 * audio_downsample_rate
|
||||
total_token_for_llm_B += (text_length + speech_length / 8) / 1000 / 1000 / 1000
|
||||
contents.append(contents_i)
|
||||
|
||||
self.contents = contents
|
||||
|
||||
logging.info(
|
||||
f"\n\ntotal_num of samplers: {len(self.contents)}, total_whrs: {total_whrs:.5f}, total_token_for_llm_B: {total_token_for_llm_B:.5g}, {path}, {file_list}\n\n")
|
||||
|
||||
def __len__(self):
|
||||
"""Internal: len ."""
|
||||
return len(self.contents)
|
||||
|
||||
def __getitem__(self, index):
|
||||
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
data = self.contents[index]
|
||||
|
||||
return data
|
||||
|
||||
def get_source_len(self, data_dict):
|
||||
"""Get source len.
|
||||
|
||||
Args:
|
||||
data_dict: TODO.
|
||||
"""
|
||||
source_len = data_dict.get("source_len", -1)
|
||||
if source_len < 0:
|
||||
source_len = len(data_dict["system"]) + len(data_dict["user"])
|
||||
return source_len
|
||||
|
||||
def get_target_len(self, data_dict):
|
||||
|
||||
"""Get target len.
|
||||
|
||||
Args:
|
||||
data_dict: TODO.
|
||||
"""
|
||||
return 0
|
||||
@@ -0,0 +1,371 @@
|
||||
import numpy as np
|
||||
from funasr.register import tables
|
||||
import logging
|
||||
import random
|
||||
import re
|
||||
|
||||
|
||||
@tables.register("prompt_classes", "MultiContextPrompt")
|
||||
class MultiContextPrompt:
|
||||
CONTEXT_TEMPLATES = {
|
||||
'en': {
|
||||
'header': "Please combine the context information provided below to complete the speech transcription task more accurately. If there is no relevant information, we will leave it blank.\n",
|
||||
'fields': {
|
||||
'hist_context': "Historical transcription: {hist_context}\n",
|
||||
'one_pass_result': "One-pass result: {one_pass_result}\n",
|
||||
'hotwords': "Hotword list: {hotwords}\n"
|
||||
}
|
||||
},
|
||||
'zh': {
|
||||
'header': "请结合下面提供的上下文信息,更加准确地完成语音转写任务。如果没有相关信息,我们会留空。\n",
|
||||
'fields': {
|
||||
'hist_context': "历史转写结果:{hist_context}\n",
|
||||
'one_pass_result': "一遍解码结果:{one_pass_result}\n",
|
||||
'hotwords': "热词列表:{hotwords}\n"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
def __init__(self,
|
||||
use_hist=True,
|
||||
use_one_pass_result=True,
|
||||
use_hotwords=True,
|
||||
use_asr_hotwords=True,
|
||||
use_multi_lingual_prompt=True,
|
||||
**kwargs):
|
||||
"""Initialize MultiContextPrompt.
|
||||
|
||||
Args:
|
||||
use_hist: TODO.
|
||||
use_one_pass_result: TODO.
|
||||
use_hotwords: TODO.
|
||||
use_asr_hotwords: TODO.
|
||||
use_multi_lingual_prompt: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
self.use_hist = use_hist
|
||||
self.use_one_pass_result = use_one_pass_result
|
||||
self.use_hotwords = use_hotwords
|
||||
self.use_asr_hotwords = use_asr_hotwords
|
||||
self.use_multi_lingual_prompt = use_multi_lingual_prompt
|
||||
self.kwargs = kwargs
|
||||
|
||||
chinese_hotwords_list = kwargs.get("chinese_hotwords_list", "")
|
||||
english_hotwords_list = kwargs.get("english_hotwords_list", "")
|
||||
if chinese_hotwords_list:
|
||||
self.chinese_hotwords_list, self.chinese_hotwords_num = self.get_hotwords_list(chinese_hotwords_list)
|
||||
else:
|
||||
self.chinese_hotwords_list = None
|
||||
self.chinese_hotwords_num = 0
|
||||
logging.info(f"chinese_hotwords_num: {self.chinese_hotwords_num}")
|
||||
|
||||
if english_hotwords_list:
|
||||
self.english_hotwords_list, self.english_hotwords_num = self.get_hotwords_list(english_hotwords_list)
|
||||
else:
|
||||
self.english_hotwords_list = None
|
||||
self.english_hotwords_num = 0
|
||||
logging.info(f"english_hotwords_num: {self.english_hotwords_num}")
|
||||
|
||||
self.max_neg_hotwords_num = kwargs.get("max_neg_hotwords_num", 900)
|
||||
self.min_neg_hotwords_num = kwargs.get("min_neg_hotwords_num", 0)
|
||||
|
||||
def get_hotwords_list(self, hotwords_file):
|
||||
"""Get hotwords list.
|
||||
|
||||
Args:
|
||||
hotwords_file: TODO.
|
||||
"""
|
||||
with open(hotwords_file, "r") as f:
|
||||
hotwords_list = f.read().strip().split("\n")
|
||||
return hotwords_list, len(hotwords_list)
|
||||
|
||||
def detect_language(self, text):
|
||||
"""Detect language.
|
||||
|
||||
Args:
|
||||
text: Text tensor or string input.
|
||||
"""
|
||||
if isinstance(text, list):
|
||||
text = " ".join(text)
|
||||
|
||||
chinese_pattern = re.compile(
|
||||
"["
|
||||
"\u4e00-\u9fff" # CJK Unified Ideographs
|
||||
"]+"
|
||||
)
|
||||
|
||||
english_pattern = re.compile(r'[A-Za-z]+')
|
||||
|
||||
chinese_matches = chinese_pattern.findall(text)
|
||||
english_matches = english_pattern.findall(text)
|
||||
|
||||
chinese_length = sum(len(match) for match in chinese_matches)
|
||||
english_length = sum(len(match) for match in english_matches)
|
||||
|
||||
total_length = len(text)
|
||||
|
||||
if total_length == 0:
|
||||
return 'zh'
|
||||
|
||||
if (chinese_length > english_length) and (chinese_length / total_length > 0.3):
|
||||
return 'zh'
|
||||
else:
|
||||
return 'en'
|
||||
|
||||
def hotwords_sampling(self, hotwords):
|
||||
|
||||
# hotwords_list = hotwords.split(", ")
|
||||
"""Hotwords sampling.
|
||||
|
||||
Args:
|
||||
hotwords: TODO.
|
||||
"""
|
||||
hotwords_list = hotwords
|
||||
selected_hotwords = []
|
||||
if self.max_neg_hotwords_num > -1:
|
||||
max_neg_hotwords_num = min(self.max_neg_hotwords_num, len(hotwords_list))
|
||||
else:
|
||||
max_neg_hotwords_num = len(hotwords_list)
|
||||
|
||||
if self.min_neg_hotwords_num < max_neg_hotwords_num:
|
||||
selected_hotwords_num = np.random.randint(self.min_neg_hotwords_num, max_neg_hotwords_num + 1)
|
||||
else:
|
||||
selected_hotwords_num = max_neg_hotwords_num
|
||||
if selected_hotwords_num > 0:
|
||||
selected_hotwords = np.random.choice(hotwords_list, selected_hotwords_num, replace=False).tolist()
|
||||
|
||||
return selected_hotwords, selected_hotwords_num
|
||||
|
||||
def get_prompt(self, item, language):
|
||||
"""Get prompt.
|
||||
|
||||
Args:
|
||||
item: TODO.
|
||||
language: Language identifier.
|
||||
"""
|
||||
template = self.CONTEXT_TEMPLATES[language]
|
||||
|
||||
prompt = template['header']
|
||||
|
||||
context_lines = []
|
||||
|
||||
if self.use_hist and item.get("hist_context"):
|
||||
context_lines.append(template['fields']['hist_context'].format(hist_context=item["hist_context"]))
|
||||
|
||||
if self.use_one_pass_result and item.get("one_pass_result"):
|
||||
context_lines.append(template['fields']['one_pass_result'].format(one_pass_result=item["one_pass_result"]))
|
||||
|
||||
hotwords = None
|
||||
if self.use_hotwords and item.get("hotwords"):
|
||||
hotwords = item["hotwords"]
|
||||
if self.use_asr_hotwords and item.get("asr_hotwords"):
|
||||
hotwords = item["asr_hotwords"]
|
||||
if hotwords is not None and hotwords != "":
|
||||
language = self.detect_language(hotwords)
|
||||
if language == 'en':
|
||||
neg_hotwords = self.english_hotwords_list
|
||||
else:
|
||||
neg_hotwords = self.chinese_hotwords_list
|
||||
if neg_hotwords is not None:
|
||||
selected_neg_hotwords, selected_neg_hotwords_num = self.hotwords_sampling(neg_hotwords)
|
||||
else:
|
||||
selected_neg_hotwords = []
|
||||
|
||||
if not isinstance(hotwords, list):
|
||||
pos_hotwords = hotwords.split(", ")
|
||||
else:
|
||||
pos_hotwords = hotwords
|
||||
hotwords = pos_hotwords + selected_neg_hotwords
|
||||
random.shuffle(hotwords)
|
||||
hotwords = ", ".join(hotwords)
|
||||
context_lines.append(template['fields']['hotwords'].format(hotwords=hotwords))
|
||||
|
||||
if context_lines:
|
||||
prompt += ''.join(context_lines)
|
||||
else:
|
||||
prompt += "\n\n\n"
|
||||
|
||||
return prompt
|
||||
|
||||
def get_inference_prompt(self, item, language="zh"):
|
||||
"""Get inference prompt.
|
||||
|
||||
Args:
|
||||
item: TODO.
|
||||
language: Language identifier.
|
||||
"""
|
||||
template = self.CONTEXT_TEMPLATES[language]
|
||||
|
||||
prompt = template['header']
|
||||
|
||||
context_lines = []
|
||||
|
||||
if self.use_hist and item.get("hist_context"):
|
||||
context_lines.append(template['fields']['hist_context'].format(hist_context=item["hist_context"]))
|
||||
|
||||
if self.use_one_pass_result and item.get("one_pass_result"):
|
||||
context_lines.append(template['fields']['one_pass_result'].format(one_pass_result=item["one_pass_result"]))
|
||||
|
||||
hotwords = None
|
||||
if self.use_hotwords and item.get("hotwords"):
|
||||
hotwords = item["hotwords"]
|
||||
if self.use_asr_hotwords and item.get("asr_hotwords"):
|
||||
hotwords = item["asr_hotwords"]
|
||||
if hotwords is not None and hotwords != "":
|
||||
print(f"hotwords: {hotwords}")
|
||||
language = self.detect_language(hotwords)
|
||||
if language == 'en':
|
||||
neg_hotwords = self.english_hotwords_list
|
||||
else:
|
||||
neg_hotwords = self.chinese_hotwords_list
|
||||
if neg_hotwords is not None:
|
||||
selected_neg_hotwords, selected_neg_hotwords_num = self.hotwords_sampling(neg_hotwords)
|
||||
else:
|
||||
selected_neg_hotwords = []
|
||||
|
||||
if not isinstance(hotwords, list):
|
||||
pos_hotwords = hotwords.split(", ")
|
||||
else:
|
||||
pos_hotwords = hotwords
|
||||
hotwords = pos_hotwords + selected_neg_hotwords
|
||||
print(f"selected_neg_hotwords_num: {selected_neg_hotwords_num}")
|
||||
random.shuffle(hotwords)
|
||||
hotwords = ", ".join(hotwords)
|
||||
context_lines.append(template['fields']['hotwords'].format(hotwords=hotwords))
|
||||
|
||||
if context_lines:
|
||||
prompt += ''.join(context_lines)
|
||||
else:
|
||||
prompt += "\n\n\n"
|
||||
|
||||
return prompt
|
||||
|
||||
|
||||
@tables.register("prompt_classes", "MultiContextPromptNew")
|
||||
class MultiContextPromptNew:
|
||||
CONTEXT_TEMPLATES = {
|
||||
'en': {
|
||||
'header': "Please combine the context information to complete the speech transcription task more accurately. If there is no relevant information, we will leave it blank.\n\n",
|
||||
'context_header': "**Context:**\n",
|
||||
'fields': {
|
||||
'hist_context': "Historical transcription: {hist_context}\n",
|
||||
'one_pass_result': "One-pass result: {one_pass_result}\n",
|
||||
'hotwords': "Hotword list: {hotwords}\n"
|
||||
}
|
||||
},
|
||||
'zh': {
|
||||
'header': "请结合上下文信息,更加准确地完成语音转写任务。如果没有相关信息,我们会留空。\n\n",
|
||||
'context_header': "**上下文信息:**\n",
|
||||
'fields': {
|
||||
'hist_context': "历史转写结果:{hist_context}\n",
|
||||
'one_pass_result': "一遍解码结果:{one_pass_result}\n",
|
||||
'hotwords': "热词列表:{hotwords}\n"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
def __init__(self,
|
||||
use_hist=True,
|
||||
use_one_pass_result=True,
|
||||
use_hotwords=True,
|
||||
use_multi_lingual_prompt=True,
|
||||
**kwargs):
|
||||
"""Initialize MultiContextPromptNew.
|
||||
|
||||
Args:
|
||||
use_hist: TODO.
|
||||
use_one_pass_result: TODO.
|
||||
use_hotwords: TODO.
|
||||
use_multi_lingual_prompt: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
self.use_hist = use_hist
|
||||
self.use_one_pass_result = use_one_pass_result
|
||||
self.use_hotwords = use_hotwords
|
||||
self.use_multi_lingual_prompt = use_multi_lingual_prompt
|
||||
|
||||
self.use_full_hotwords_ratio = kwargs.get("use_full_hotwords_ratio", 0.2)
|
||||
self.max_hotwords_num = kwargs.get("max_hotwords_num", -1)
|
||||
self.min_hotwords_num = kwargs.get("min_hotwords_num", 15)
|
||||
|
||||
def hotwords_sampling(self, hotwords):
|
||||
|
||||
"""Hotwords sampling.
|
||||
|
||||
Args:
|
||||
hotwords: TODO.
|
||||
"""
|
||||
hotwords_list = hotwords.split(", ")
|
||||
if self.max_hotwords_num > 0:
|
||||
max_hotwords_num = min(self.max_hotwords_num, len(hotwords_list))
|
||||
else:
|
||||
max_hotwords_num = len(hotwords_list)
|
||||
|
||||
if self.min_hotwords_num < max_hotwords_num:
|
||||
selected_hotwords_num = np.random.randint(self.min_hotwords_num, max_hotwords_num + 1)
|
||||
else:
|
||||
selected_hotwords_num = max_hotwords_num
|
||||
|
||||
selected_hotwords = np.random.choice(hotwords_list, selected_hotwords_num, replace=False)
|
||||
hotwords_list = ", ".join(selected_hotwords)
|
||||
|
||||
return hotwords_list, selected_hotwords_num
|
||||
|
||||
def get_prompt(self, item, language):
|
||||
"""Get prompt.
|
||||
|
||||
Args:
|
||||
item: TODO.
|
||||
language: Language identifier.
|
||||
"""
|
||||
template = self.CONTEXT_TEMPLATES[language]
|
||||
|
||||
prompt = template['header']
|
||||
|
||||
context_lines = []
|
||||
|
||||
if self.use_hist and item.get("hist_context"):
|
||||
context_lines.append(template['fields']['hist_context'].format(hist_context=item["hist_context"]))
|
||||
|
||||
if self.use_one_pass_result and item.get("one_pass_result"):
|
||||
context_lines.append(template['fields']['one_pass_result'].format(one_pass_result=item["one_pass_result"]))
|
||||
|
||||
if self.use_hotwords and item.get("hotwords"):
|
||||
hotwords = item["hotwords"]
|
||||
if np.random.rand() < self.use_full_hotwords_ratio:
|
||||
hotwords = hotwords
|
||||
else:
|
||||
hotwords, selected_hotwords_num = self.hotwords_sampling(hotwords)
|
||||
context_lines.append(template['fields']['hotwords'].format(hotwords=hotwords))
|
||||
|
||||
if context_lines:
|
||||
prompt += template['context_header'] + ''.join(context_lines)
|
||||
|
||||
return prompt
|
||||
|
||||
def get_inference_prompt(self, hist_context="", one_pass_result="", hotwords=""):
|
||||
"""Get inference prompt.
|
||||
|
||||
Args:
|
||||
hist_context: TODO.
|
||||
one_pass_result: TODO.
|
||||
hotwords: TODO.
|
||||
"""
|
||||
language = 'zh' if self.use_multi_lingual_prompt and np.random.rand() < 0.5 else 'en'
|
||||
template = self.CONTEXT_TEMPLATES[language]
|
||||
|
||||
prompt = template['header']
|
||||
|
||||
context_lines = []
|
||||
|
||||
if hist_context:
|
||||
context_lines.append(template['fields']['hist_context'].format(hist_context=hist_context))
|
||||
if one_pass_result:
|
||||
context_lines.append(template['fields']['one_pass_result'].format(one_pass_result=one_pass_result))
|
||||
if hotwords:
|
||||
context_lines.append(template['fields']['hotwords'].format(hotwords=hotwords))
|
||||
|
||||
if context_lines:
|
||||
prompt += template['context_header'] + ''.join(context_lines)
|
||||
|
||||
return prompt
|
||||
@@ -0,0 +1,165 @@
|
||||
import torch
|
||||
import random
|
||||
|
||||
|
||||
from funasr.register import tables
|
||||
from funasr.utils.load_utils import extract_fbank, load_audio_text_image_video
|
||||
|
||||
|
||||
@tables.register("dataset_classes", "KwsMTDataset")
|
||||
class KwsMTDataset(torch.utils.data.Dataset):
|
||||
"""
|
||||
KwsMTDataset, support multi tokenizers
|
||||
"""
|
||||
def __init__(self,
|
||||
path,
|
||||
index_ds: str = None,
|
||||
frontend=None,
|
||||
tokenizer=None,
|
||||
is_training: bool = True,
|
||||
int_pad_value: int = -1,
|
||||
float_pad_value: float = 0.0,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize KwsMTDataset.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
index_ds: TODO.
|
||||
frontend: Audio frontend for feature extraction.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
is_training: Boolean flag for training.
|
||||
int_pad_value: TODO.
|
||||
float_pad_value: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
index_ds_class = tables.index_ds_classes.get(index_ds)
|
||||
self.index_ds = index_ds_class(path, **kwargs)
|
||||
|
||||
self.preprocessor_speech = None
|
||||
self.preprocessor_text = None
|
||||
|
||||
if is_training:
|
||||
preprocessor_speech = kwargs.get("preprocessor_speech", None)
|
||||
if preprocessor_speech:
|
||||
preprocessor_speech_class = tables.preprocessor_classes.get(preprocessor_speech)
|
||||
preprocessor_speech = preprocessor_speech_class(
|
||||
**kwargs.get("preprocessor_speech_conf")
|
||||
)
|
||||
self.preprocessor_speech = preprocessor_speech
|
||||
|
||||
preprocessor_text = kwargs.get("preprocessor_text", None)
|
||||
if preprocessor_text:
|
||||
preprocessor_text_class = tables.preprocessor_classes.get(preprocessor_text)
|
||||
preprocessor_text = preprocessor_text_class(**kwargs.get("preprocessor_text_conf"))
|
||||
self.preprocessor_text = preprocessor_text
|
||||
|
||||
self.frontend = frontend
|
||||
self.fs = 16000 if frontend is None else frontend.fs
|
||||
self.data_type = "sound"
|
||||
print(tokenizer)
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.int_pad_value = int_pad_value
|
||||
self.float_pad_value = float_pad_value
|
||||
|
||||
def get_source_len(self, index):
|
||||
"""Get source len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_source_len(item)
|
||||
|
||||
def get_target_len(self, index):
|
||||
"""Get target len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_target_len(item)
|
||||
|
||||
def __len__(self):
|
||||
"""Internal: len ."""
|
||||
return len(self.index_ds)
|
||||
|
||||
def __getitem__(self, index):
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
# import pdb;
|
||||
# pdb.set_trace()
|
||||
source = item["source"]
|
||||
data_src = load_audio_text_image_video(source, fs=self.fs)
|
||||
if self.preprocessor_speech:
|
||||
data_src = self.preprocessor_speech(data_src, fs=self.fs)
|
||||
speech, speech_lengths = extract_fbank(
|
||||
data_src, data_type=self.data_type, frontend=self.frontend, is_final=True
|
||||
) # speech: [b, T, d]
|
||||
|
||||
target = item["target"]
|
||||
if self.preprocessor_text:
|
||||
target = self.preprocessor_text(target)
|
||||
|
||||
if self.tokenizer[0]:
|
||||
ids = self.tokenizer[0].encode(target)
|
||||
text = torch.tensor(ids, dtype=torch.int64)
|
||||
# print("target: ", target, ", ids: ", str(ids))
|
||||
else:
|
||||
ids = target
|
||||
text = ids
|
||||
|
||||
if self.tokenizer[1]:
|
||||
ids2 = self.tokenizer[1].encode(target)
|
||||
text2 = torch.tensor(ids2, dtype=torch.int64)
|
||||
# print("target: ", target, ", ids2: ", str(ids2))
|
||||
else:
|
||||
ids2 = target
|
||||
text2 = ids2
|
||||
|
||||
ids_lengths = len(ids)
|
||||
text_lengths = torch.tensor([ids_lengths], dtype=torch.int32)
|
||||
|
||||
ids2_lengths = len(ids2)
|
||||
text2_lengths = torch.tensor([ids2_lengths], dtype=torch.int32)
|
||||
|
||||
return {"speech": speech[0, :, :],
|
||||
"speech_lengths": speech_lengths,
|
||||
"text": text,
|
||||
"text_lengths": text_lengths,
|
||||
"text2": text2,
|
||||
"text2_lengths": text2_lengths,
|
||||
}
|
||||
|
||||
|
||||
def collator(self, samples: list=None):
|
||||
"""Collator.
|
||||
|
||||
Args:
|
||||
samples: TODO.
|
||||
"""
|
||||
outputs = {}
|
||||
for sample in samples:
|
||||
for key in sample.keys():
|
||||
if key not in outputs:
|
||||
outputs[key] = []
|
||||
outputs[key].append(sample[key])
|
||||
|
||||
for key, data_list in outputs.items():
|
||||
if isinstance(data_list[0], torch.Tensor):
|
||||
if data_list[0].dtype == torch.int64 or data_list[0].dtype == torch.int32:
|
||||
|
||||
pad_value = self.int_pad_value
|
||||
else:
|
||||
pad_value = self.float_pad_value
|
||||
|
||||
outputs[key] = torch.nn.utils.rnn.pad_sequence(
|
||||
data_list, batch_first=True, padding_value=pad_value
|
||||
)
|
||||
return outputs
|
||||
@@ -0,0 +1,530 @@
|
||||
import torch
|
||||
import copy
|
||||
|
||||
from funasr.register import tables
|
||||
from funasr.utils.load_utils import extract_fbank, load_audio_text_image_video
|
||||
|
||||
|
||||
@tables.register("dataset_classes", "AudioLLMNARDataset")
|
||||
class AudioLLMNARDataset(torch.utils.data.Dataset):
|
||||
"""
|
||||
AudioLLMDataset
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path,
|
||||
index_ds: str = None,
|
||||
frontend=None,
|
||||
tokenizer=None,
|
||||
int_pad_value: int = -1,
|
||||
float_pad_value: float = 0.0,
|
||||
**kwargs
|
||||
):
|
||||
"""Initialize AudioLLMNARDataset.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
index_ds: TODO.
|
||||
frontend: Audio frontend for feature extraction.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
int_pad_value: TODO.
|
||||
float_pad_value: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
index_ds_class = tables.index_ds_classes.get(index_ds)
|
||||
self.index_ds = index_ds_class(path, **kwargs)
|
||||
preprocessor_speech = kwargs.get("preprocessor_speech", None)
|
||||
if preprocessor_speech:
|
||||
preprocessor_speech_class = tables.preprocessor_classes.get(preprocessor_speech)
|
||||
preprocessor_speech = preprocessor_speech_class(
|
||||
**kwargs.get("preprocessor_speech_conf", {})
|
||||
)
|
||||
self.preprocessor_speech = preprocessor_speech
|
||||
preprocessor_text = kwargs.get("preprocessor_text", None)
|
||||
if preprocessor_text:
|
||||
preprocessor_text_class = tables.preprocessor_classes.get(preprocessor_text)
|
||||
preprocessor_text = preprocessor_text_class(**kwargs.get("preprocessor_text_conf", {}))
|
||||
self.preprocessor_text = preprocessor_text
|
||||
|
||||
self.frontend = frontend
|
||||
self.fs = 16000 if frontend is None else frontend.fs
|
||||
self.data_type = "sound"
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.float_pad_value = float_pad_value
|
||||
self.prompt = kwargs.get("prompt", "Please copy the following text.")
|
||||
self.prompt_pre = "USER: \nINSTRUCTION: {}\nINPUT: ".format(
|
||||
self.prompt
|
||||
) # "USER: \nINSTRUCTION: {}\nINPUT: {}\nASSISTANT: "
|
||||
self.prompt_af = ""
|
||||
self.IGNORE_INDEX = kwargs.get("IGNORE_INDEX", -100)
|
||||
self.int_pad_value = self.IGNORE_INDEX
|
||||
|
||||
def get_source_len(self, index):
|
||||
"""Get source len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_source_len(item)
|
||||
|
||||
def get_target_len(self, index):
|
||||
"""Get target len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_target_len(item)
|
||||
|
||||
def __len__(self):
|
||||
"""Internal: len ."""
|
||||
return len(self.index_ds)
|
||||
|
||||
def __getitem__(self, index):
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
source = item["source"]
|
||||
data_src = load_audio_text_image_video(source, fs=self.fs)
|
||||
if self.preprocessor_speech:
|
||||
data_src = self.preprocessor_speech(data_src, fs=self.fs)
|
||||
speech, speech_lengths = extract_fbank(
|
||||
data_src, data_type=self.data_type, frontend=self.frontend, is_final=True
|
||||
) # speech: [b, T, d]
|
||||
speech = speech.squeeze(0)
|
||||
|
||||
target = item["target"]
|
||||
if self.preprocessor_text:
|
||||
target = self.preprocessor_text(target)
|
||||
|
||||
prompt_ids_pre = self.tokenizer.encode(self.prompt_pre) # [bos,prompt]
|
||||
prompt_ids_length = len(prompt_ids_pre)
|
||||
|
||||
# bos prompt audio bos target
|
||||
# prompt_input = "{}{}".format(self.prompt_pre, target)
|
||||
# prompt_input_ids = self.tokenizer.encode(prompt_input) #[bos, prompt, input]
|
||||
# audio_length = len(prompt_input_ids) - prompt_ids_length
|
||||
target_ids = self.tokenizer.encode(target)
|
||||
if target_ids[0] == self.tokenizer.bos_token_id:
|
||||
target_ids = target_ids[1:]
|
||||
target_ids_length = len(target_ids)
|
||||
audio_length = target_ids_length
|
||||
input_ids = (
|
||||
prompt_ids_pre + target_ids + [self.tokenizer.pad_token_id] + target_ids
|
||||
) # [bos, prompt, input, pad, target]
|
||||
input_ids = torch.tensor(
|
||||
copy.deepcopy(input_ids), dtype=torch.int64
|
||||
) # [bos, prompt, input, pad, target]
|
||||
input_ids[prompt_ids_length : prompt_ids_length + audio_length] = (
|
||||
-1
|
||||
) # [bos, prompt,-1, pad, target] # it is no need, only for check
|
||||
attention_mask = input_ids.ge(-1) # [true, true, true, true, true], length mask
|
||||
|
||||
# bos prompt audio target eos
|
||||
# prompt_answer = "{}{}".format(self.prompt_pre, target)
|
||||
# prompt_answer_ids = self.tokenizer.encode(prompt_answer) #[bos, prompt, input]
|
||||
# answer_length = len(prompt_answer_ids) - prompt_ids_length
|
||||
target_ids = self.tokenizer.encode(target)
|
||||
if target_ids[0] == self.tokenizer.bos_token_id:
|
||||
target_ids = target_ids[1:]
|
||||
# target_ids_length = len(target_ids)
|
||||
labels_ids = (
|
||||
prompt_ids_pre + target_ids + target_ids + [self.tokenizer.eos_token_id]
|
||||
) # [bos, prompt, input, target, eos]
|
||||
labels_ids = torch.tensor(
|
||||
copy.deepcopy(labels_ids), dtype=torch.int64
|
||||
) # [bos, prompt, input, target, eos]
|
||||
labels_ids[:prompt_ids_length] = -1 # [-1, -1, input, target, eos]
|
||||
label_mask = labels_ids.ge(0) # [false, false, true, true, true], length mask
|
||||
labels_ids[~label_mask] = self.IGNORE_INDEX # [-1, -1, input, target, eos]
|
||||
|
||||
audio_mask = (
|
||||
[0] * prompt_ids_length + [1] * audio_length + [0] * target_ids_length + [0]
|
||||
) # [0, 0, 1, 0, 0]
|
||||
audio_mask = torch.tensor(audio_mask, dtype=torch.float32)
|
||||
|
||||
ids = target_ids # self.tokenizer.encode(target) # token ids is different from labels_ids
|
||||
text = torch.tensor(ids, dtype=torch.int64)
|
||||
text_lengths = torch.tensor([len(ids)], dtype=torch.int32)
|
||||
|
||||
prompt_bos_length = torch.tensor([len(prompt_ids_pre)], dtype=torch.int32)
|
||||
|
||||
return {
|
||||
"speech": speech,
|
||||
"speech_lengths": speech_lengths,
|
||||
"text": text,
|
||||
"text_lengths": text_lengths,
|
||||
"input_ids": input_ids,
|
||||
"attention_mask": attention_mask,
|
||||
"labels_ids": labels_ids,
|
||||
"label_mask": label_mask,
|
||||
"audio_mask": audio_mask,
|
||||
"prompt_bos_length": prompt_bos_length,
|
||||
}
|
||||
|
||||
def collator(self, samples: list = None):
|
||||
"""Collator.
|
||||
|
||||
Args:
|
||||
samples: TODO.
|
||||
"""
|
||||
outputs = {}
|
||||
for sample in samples:
|
||||
for key in sample.keys():
|
||||
if key not in outputs:
|
||||
outputs[key] = []
|
||||
outputs[key].append(sample[key])
|
||||
|
||||
for key, data_list in outputs.items():
|
||||
if isinstance(data_list[0], torch.Tensor):
|
||||
if data_list[0].dtype == torch.int64 or data_list[0].dtype == torch.int32:
|
||||
|
||||
pad_value = self.int_pad_value
|
||||
else:
|
||||
pad_value = self.float_pad_value
|
||||
|
||||
outputs[key] = torch.nn.utils.rnn.pad_sequence(
|
||||
data_list, batch_first=True, padding_value=pad_value
|
||||
)
|
||||
return outputs
|
||||
|
||||
|
||||
@tables.register("dataset_classes", "AudioLLMDataset")
|
||||
class AudioLLMDataset(torch.utils.data.Dataset):
|
||||
"""
|
||||
AudioLLMDataset
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path,
|
||||
index_ds: str = None,
|
||||
frontend=None,
|
||||
tokenizer=None,
|
||||
int_pad_value: int = -1,
|
||||
float_pad_value: float = 0.0,
|
||||
**kwargs
|
||||
):
|
||||
"""Initialize AudioLLMDataset.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
index_ds: TODO.
|
||||
frontend: Audio frontend for feature extraction.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
int_pad_value: TODO.
|
||||
float_pad_value: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
index_ds_class = tables.index_ds_classes.get(index_ds)
|
||||
self.index_ds = index_ds_class(path, **kwargs)
|
||||
preprocessor_speech = kwargs.get("preprocessor_speech", None)
|
||||
if preprocessor_speech:
|
||||
preprocessor_speech_class = tables.preprocessor_classes.get(preprocessor_speech)
|
||||
preprocessor_speech = preprocessor_speech_class(
|
||||
**kwargs.get("preprocessor_speech_conf", {})
|
||||
)
|
||||
self.preprocessor_speech = preprocessor_speech
|
||||
preprocessor_text = kwargs.get("preprocessor_text", None)
|
||||
if preprocessor_text:
|
||||
preprocessor_text_class = tables.preprocessor_classes.get(preprocessor_text)
|
||||
preprocessor_text = preprocessor_text_class(**kwargs.get("preprocessor_text_conf", {}))
|
||||
self.preprocessor_text = preprocessor_text
|
||||
|
||||
self.frontend = frontend
|
||||
self.fs = 16000 if frontend is None else frontend.fs
|
||||
self.data_type = "sound"
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.float_pad_value = float_pad_value
|
||||
self.prompt = kwargs.get("prompt", "Transcribe speech to text.")
|
||||
self.prompt_pre = "USER: \nINSTRUCTION: {}\nINPUT: ".format(
|
||||
self.prompt
|
||||
) # "USER: \nINSTRUCTION: {}\nnINPUT: {}\nASSISTANT: "
|
||||
self.prompt_af = ""
|
||||
self.IGNORE_INDEX = kwargs.get("IGNORE_INDEX", -100)
|
||||
self.int_pad_value = self.IGNORE_INDEX
|
||||
|
||||
def get_source_len(self, index):
|
||||
"""Get source len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_source_len(item)
|
||||
|
||||
def get_target_len(self, index):
|
||||
"""Get target len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_target_len(item)
|
||||
|
||||
def __len__(self):
|
||||
"""Internal: len ."""
|
||||
return len(self.index_ds)
|
||||
|
||||
def __getitem__(self, index):
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
# import pdb;
|
||||
# pdb.set_trace()
|
||||
source = item["source"]
|
||||
data_src = load_audio_text_image_video(source, fs=self.fs)
|
||||
if self.preprocessor_speech:
|
||||
data_src = self.preprocessor_speech(data_src, fs=self.fs)
|
||||
speech, speech_lengths = extract_fbank(
|
||||
data_src, data_type=self.data_type, frontend=self.frontend, is_final=True
|
||||
) # speech: [b, T, d]
|
||||
speech = speech.squeeze(0)
|
||||
|
||||
target = item["target"]
|
||||
if self.preprocessor_text:
|
||||
target = self.preprocessor_text(target)
|
||||
|
||||
prompt_ids_pre = self.tokenizer.encode(self.prompt_pre) # [bos,prompt]
|
||||
prompt_ids_length = len(prompt_ids_pre)
|
||||
|
||||
prompt_input = "{}{}".format(self.prompt_pre, target)
|
||||
prompt_input_ids = self.tokenizer.encode(prompt_input)
|
||||
audio_length = len(prompt_input_ids) - prompt_ids_length
|
||||
input_ids = prompt_input_ids + [self.tokenizer.pad_token_id]
|
||||
input_ids = torch.tensor(input_ids, dtype=torch.int64) # [bos, prompt, input, pad]
|
||||
input_ids[prompt_ids_length:] = -1 # [bos, prompt,-1,-1]
|
||||
attention_mask = input_ids.ge(-1) # [true, true, true, true], length mask
|
||||
|
||||
prompt_answer = "{}{}".format(self.prompt_pre, target)
|
||||
prompt_answer_ids = self.tokenizer.encode(prompt_answer)
|
||||
answer_length = len(prompt_answer_ids) - prompt_ids_length
|
||||
labels_ids = copy.deepcopy(prompt_input_ids) + [self.tokenizer.eos_token_id]
|
||||
labels_ids = torch.tensor(labels_ids, dtype=torch.int64) # [bos, prompt, input, eos]
|
||||
labels_ids[:prompt_ids_length] = -1 # [-1, -1, input, eos]
|
||||
label_mask = labels_ids.ge(0) # [False,False,True,True]
|
||||
labels_ids[~label_mask] = self.IGNORE_INDEX # [-100,-100,input,eos]
|
||||
|
||||
audio_mask = [0] * prompt_ids_length + [1] * audio_length + [0]
|
||||
audio_mask = torch.tensor(audio_mask, dtype=torch.float32)
|
||||
|
||||
ids = self.tokenizer.encode(target) # token ids is different from labels_ids
|
||||
text = torch.tensor(ids, dtype=torch.int64)
|
||||
text_lengths = torch.tensor([len(ids)], dtype=torch.int32)
|
||||
|
||||
return {
|
||||
"speech": speech,
|
||||
"speech_lengths": speech_lengths,
|
||||
"text": text,
|
||||
"text_lengths": text_lengths,
|
||||
"input_ids": input_ids,
|
||||
"attention_mask": attention_mask,
|
||||
"labels_ids": labels_ids,
|
||||
"label_mask": label_mask,
|
||||
"audio_mask": audio_mask,
|
||||
}
|
||||
|
||||
def collator(self, samples: list = None):
|
||||
"""Collator.
|
||||
|
||||
Args:
|
||||
samples: TODO.
|
||||
"""
|
||||
outputs = {}
|
||||
for sample in samples:
|
||||
for key in sample.keys():
|
||||
if key not in outputs:
|
||||
outputs[key] = []
|
||||
outputs[key].append(sample[key])
|
||||
|
||||
for key, data_list in outputs.items():
|
||||
if isinstance(data_list[0], torch.Tensor):
|
||||
if data_list[0].dtype == torch.int64 or data_list[0].dtype == torch.int32:
|
||||
|
||||
pad_value = self.int_pad_value
|
||||
else:
|
||||
pad_value = self.float_pad_value
|
||||
|
||||
outputs[key] = torch.nn.utils.rnn.pad_sequence(
|
||||
data_list, batch_first=True, padding_value=pad_value
|
||||
)
|
||||
return outputs
|
||||
|
||||
|
||||
@tables.register("dataset_classes", "AudioLLMARDataset")
|
||||
class AudioLLMARDataset(torch.utils.data.Dataset):
|
||||
"""
|
||||
AudioLLMDataset
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path,
|
||||
index_ds: str = None,
|
||||
frontend=None,
|
||||
tokenizer=None,
|
||||
int_pad_value: int = -1,
|
||||
float_pad_value: float = 0.0,
|
||||
**kwargs
|
||||
):
|
||||
"""Initialize AudioLLMARDataset.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
index_ds: TODO.
|
||||
frontend: Audio frontend for feature extraction.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
int_pad_value: TODO.
|
||||
float_pad_value: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
index_ds_class = tables.index_ds_classes.get(index_ds)
|
||||
self.index_ds = index_ds_class(path, **kwargs)
|
||||
preprocessor_speech = kwargs.get("preprocessor_speech", None)
|
||||
if preprocessor_speech:
|
||||
preprocessor_speech_class = tables.preprocessor_classes.get(preprocessor_speech)
|
||||
preprocessor_speech = preprocessor_speech_class(
|
||||
**kwargs.get("preprocessor_speech_conf", {})
|
||||
)
|
||||
self.preprocessor_speech = preprocessor_speech
|
||||
preprocessor_text = kwargs.get("preprocessor_text", None)
|
||||
if preprocessor_text:
|
||||
preprocessor_text_class = tables.preprocessor_classes.get(preprocessor_text)
|
||||
preprocessor_text = preprocessor_text_class(**kwargs.get("preprocessor_text_conf", {}))
|
||||
self.preprocessor_text = preprocessor_text
|
||||
|
||||
self.frontend = frontend
|
||||
self.fs = 16000 if frontend is None else frontend.fs
|
||||
self.data_type = "sound"
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.float_pad_value = float_pad_value
|
||||
self.prompt = kwargs.get("prompt", "Transcribe speech to text.")
|
||||
self.prompt_pre = "USER: \nINSTRUCTION: {}\nINPUT: ".format(
|
||||
self.prompt
|
||||
) # "USER: \nINSTRUCTION: {}\nnINPUT: {}\nASSISTANT: "
|
||||
self.prompt_af = ""
|
||||
self.IGNORE_INDEX = kwargs.get("IGNORE_INDEX", -100)
|
||||
self.int_pad_value = self.IGNORE_INDEX
|
||||
|
||||
def get_source_len(self, index):
|
||||
"""Get source len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_source_len(item)
|
||||
|
||||
def get_target_len(self, index):
|
||||
"""Get target len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_target_len(item)
|
||||
|
||||
def __len__(self):
|
||||
"""Internal: len ."""
|
||||
return len(self.index_ds)
|
||||
|
||||
def __getitem__(self, index):
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
# import pdb;
|
||||
# pdb.set_trace()
|
||||
source = item["source"]
|
||||
data_src = load_audio_text_image_video(source, fs=self.fs)
|
||||
if self.preprocessor_speech:
|
||||
data_src = self.preprocessor_speech(data_src, fs=self.fs)
|
||||
speech, speech_lengths = extract_fbank(
|
||||
data_src, data_type=self.data_type, frontend=self.frontend, is_final=True
|
||||
) # speech: [b, T, d]
|
||||
speech = speech.squeeze(0)
|
||||
|
||||
target = item["target"]
|
||||
if self.preprocessor_text:
|
||||
target = self.preprocessor_text(target)
|
||||
|
||||
prompt_ids_pre = self.tokenizer.encode(self.prompt_pre) # [bos,prompt]
|
||||
prompt_ids_length = len(prompt_ids_pre)
|
||||
|
||||
prompt_input = "{}{}".format(self.prompt_pre, target)
|
||||
prompt_input_ids = self.tokenizer.encode(prompt_input)
|
||||
audio_length = len(prompt_input_ids) - prompt_ids_length
|
||||
input_ids = prompt_input_ids + [self.tokenizer.pad_token_id]
|
||||
input_ids = torch.tensor(input_ids, dtype=torch.int64) # [bos, prompt, input, pad]
|
||||
input_ids[prompt_ids_length:] = -1 # [bos, prompt,-1,-1]
|
||||
attention_mask = input_ids.ge(-1) # [true, true, true, true], length mask
|
||||
|
||||
prompt_answer = "{}{}".format(self.prompt_pre, target)
|
||||
prompt_answer_ids = self.tokenizer.encode(prompt_answer)
|
||||
answer_length = len(prompt_answer_ids) - prompt_ids_length
|
||||
labels_ids = copy.deepcopy(prompt_input_ids) + [self.tokenizer.eos_token_id]
|
||||
labels_ids = torch.tensor(labels_ids, dtype=torch.int64) # [bos, prompt, input, eos]
|
||||
labels_ids[:prompt_ids_length] = -1 # [-1, -1, input, eos]
|
||||
label_mask = labels_ids.ge(0) # [False,False,True,True]
|
||||
labels_ids[~label_mask] = self.IGNORE_INDEX # [-100,-100,input,eos]
|
||||
|
||||
audio_mask = [0] * prompt_ids_length + [1] * audio_length + [0]
|
||||
audio_mask = torch.tensor(audio_mask, dtype=torch.float32)
|
||||
|
||||
ids = self.tokenizer.encode(target) # token ids is different from labels_ids
|
||||
text = torch.tensor(ids, dtype=torch.int64)
|
||||
text_lengths = torch.tensor([len(ids)], dtype=torch.int32)
|
||||
|
||||
return {
|
||||
"speech": speech,
|
||||
"speech_lengths": speech_lengths,
|
||||
"text": text,
|
||||
"text_lengths": text_lengths,
|
||||
"input_ids": input_ids,
|
||||
"attention_mask": attention_mask,
|
||||
"labels_ids": labels_ids,
|
||||
"label_mask": label_mask,
|
||||
"audio_mask": audio_mask,
|
||||
}
|
||||
|
||||
def collator(self, samples: list = None):
|
||||
"""Collator.
|
||||
|
||||
Args:
|
||||
samples: TODO.
|
||||
"""
|
||||
outputs = {}
|
||||
for sample in samples:
|
||||
for key in sample.keys():
|
||||
if key not in outputs:
|
||||
outputs[key] = []
|
||||
outputs[key].append(sample[key])
|
||||
|
||||
for key, data_list in outputs.items():
|
||||
if isinstance(data_list[0], torch.Tensor):
|
||||
if data_list[0].dtype == torch.int64 or data_list[0].dtype == torch.int32:
|
||||
|
||||
pad_value = self.int_pad_value
|
||||
else:
|
||||
pad_value = self.float_pad_value
|
||||
|
||||
outputs[key] = torch.nn.utils.rnn.pad_sequence(
|
||||
data_list, batch_first=True, padding_value=pad_value
|
||||
)
|
||||
return outputs
|
||||
@@ -0,0 +1,45 @@
|
||||
import os
|
||||
import json
|
||||
import torch
|
||||
import logging
|
||||
import concurrent.futures
|
||||
import librosa
|
||||
import torch.distributed as dist
|
||||
from typing import Collection
|
||||
import torch
|
||||
import torchaudio
|
||||
from torch import nn
|
||||
import random
|
||||
import re
|
||||
import string
|
||||
from funasr.tokenizer.cleaner import TextCleaner
|
||||
from funasr.register import tables
|
||||
|
||||
|
||||
@tables.register("preprocessor_classes", "TextPreprocessRemovePunctuation")
|
||||
class TextPreprocessRemovePunctuation(nn.Module):
|
||||
def __init__(self, **kwargs):
|
||||
"""Initialize TextPreprocessRemovePunctuation.
|
||||
|
||||
Args:
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
def forward(self, text, **kwargs):
|
||||
# 定义英文标点符号
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
text: Text tensor or string input.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
en_punct = string.punctuation
|
||||
# 定义中文标点符号(部分常用的)
|
||||
cn_punct = "。?!,、;:“”‘’()《》【】…—~·"
|
||||
# 合并英文和中文标点符号
|
||||
all_punct = en_punct + cn_punct
|
||||
# 创建正则表达式模式,匹配任何在all_punct中的字符
|
||||
punct_pattern = re.compile("[{}]".format(re.escape(all_punct)))
|
||||
# 使用正则表达式的sub方法替换掉这些字符
|
||||
return punct_pattern.sub("", text)
|
||||
@@ -0,0 +1,191 @@
|
||||
import torch
|
||||
import copy
|
||||
|
||||
from funasr.register import tables
|
||||
from funasr.utils.load_utils import extract_fbank, load_audio_text_image_video
|
||||
|
||||
|
||||
@tables.register("dataset_classes", "AudioLLMQwenAudioDataset")
|
||||
class AudioLLMQwenAudioDataset(torch.utils.data.Dataset):
|
||||
"""
|
||||
AudioLLMDataset
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path,
|
||||
index_ds: str = None,
|
||||
frontend=None,
|
||||
tokenizer=None,
|
||||
int_pad_value: int = -1,
|
||||
float_pad_value: float = 0.0,
|
||||
**kwargs
|
||||
):
|
||||
"""Initialize AudioLLMQwenAudioDataset.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
index_ds: TODO.
|
||||
frontend: Audio frontend for feature extraction.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
int_pad_value: TODO.
|
||||
float_pad_value: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
index_ds_class = tables.index_ds_classes.get(index_ds)
|
||||
self.index_ds = index_ds_class(path, **kwargs)
|
||||
preprocessor_speech = kwargs.get("preprocessor_speech", None)
|
||||
if preprocessor_speech:
|
||||
preprocessor_speech_class = tables.preprocessor_classes.get(preprocessor_speech)
|
||||
preprocessor_speech = preprocessor_speech_class(
|
||||
**kwargs.get("preprocessor_speech_conf", {})
|
||||
)
|
||||
self.preprocessor_speech = preprocessor_speech
|
||||
preprocessor_text = kwargs.get("preprocessor_text", None)
|
||||
if preprocessor_text:
|
||||
preprocessor_text_class = tables.preprocessor_classes.get(preprocessor_text)
|
||||
preprocessor_text = preprocessor_text_class(**kwargs.get("preprocessor_text_conf", {}))
|
||||
self.preprocessor_text = preprocessor_text
|
||||
|
||||
self.frontend = frontend
|
||||
self.fs = 16000 if frontend is None else frontend.fs
|
||||
self.data_type = "sound"
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.float_pad_value = float_pad_value
|
||||
self.prompt = kwargs.get("prompt", "Transcribe speech to text.")
|
||||
# self.prompt_pre = "USER: \nINSTRUCTION: {}\nINPUT: ".format(self.prompt) # "USER: \nINSTRUCTION: {}\nnINPUT: {}\nASSISTANT: "
|
||||
self.prompt_af = ""
|
||||
self.IGNORE_INDEX = kwargs.get("IGNORE_INDEX", -100)
|
||||
self.int_pad_value = self.IGNORE_INDEX
|
||||
self.audio_adaptor_downsample_rate = kwargs.get("audio_adaptor_downsample_rate", 5)
|
||||
self.audio_encoder_downsample_rate = kwargs.get("audio_encoder_downsample_rate", 2)
|
||||
self.prompt_template = "{}"
|
||||
self.answer_template = "{}"
|
||||
|
||||
def get_source_len(self, index):
|
||||
"""Get source len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_source_len(item)
|
||||
|
||||
def get_target_len(self, index):
|
||||
"""Get target len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_target_len(item)
|
||||
|
||||
def __len__(self):
|
||||
"""Internal: len ."""
|
||||
return len(self.index_ds)
|
||||
|
||||
def __getitem__(self, index):
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
source = item["source"]
|
||||
data_src = load_audio_text_image_video(source, fs=self.fs)
|
||||
if self.preprocessor_speech:
|
||||
data_src = self.preprocessor_speech(data_src, fs=self.fs)
|
||||
speech, speech_lengths = extract_fbank(
|
||||
data_src, data_type=self.data_type, frontend=self.frontend, is_final=True
|
||||
) # speech: [b, T, d]
|
||||
speech = speech.squeeze(0)
|
||||
|
||||
audio_pseudo_length = (
|
||||
(speech.shape[0] + 1)
|
||||
// self.audio_adaptor_downsample_rate
|
||||
// self.audio_encoder_downsample_rate
|
||||
)
|
||||
audio_pseudo = torch.full((audio_pseudo_length,), -1) # placeholder
|
||||
|
||||
target = item["target"]
|
||||
if self.preprocessor_text:
|
||||
target = self.preprocessor_text(target)
|
||||
|
||||
self.prompt_pre = self.prompt_template.format(self.prompt)
|
||||
prompt_ids_pre = self.tokenizer.encode(self.prompt_pre) # [bos,prompt]
|
||||
prompt_pre_length = len(prompt_ids_pre)
|
||||
|
||||
# input
|
||||
input = self.answer_template.format(target.lower())
|
||||
prompt_input = "{}{}".format(self.prompt_pre, input)
|
||||
prompt_input_ids = self.tokenizer.encode(prompt_input) # [bos, prompt, input]
|
||||
# audio_length = len(prompt_input_ids) - prompt_pre_length
|
||||
input_ids = prompt_input_ids + [self.tokenizer.pad_token_id] # [bos, prompt, input, pad]
|
||||
input_ids_length = len(input_ids)
|
||||
input_ids = torch.tensor(input_ids, dtype=torch.int64) # [bos, prompt, input, pad]
|
||||
input_ids = torch.cat((audio_pseudo, input_ids)) # [audio, bos, prompt, input, pad]
|
||||
# input_ids[:audio_pseudo_length] = -1 # [-1, bos, prompt, input, pad]
|
||||
attention_mask = input_ids.ge(-1) # [true, true, true, true, true], length mask
|
||||
# input_ids[prompt_pre_length:] = -1 # [bos, prompt,-1,-1]
|
||||
# attention_mask = input_ids.ge(-1) # [true, true, true, true], length mask
|
||||
|
||||
# label
|
||||
answer = self.answer_template.format(target.lower())
|
||||
prompt_answer = "{}{}".format(self.prompt_pre, answer)
|
||||
prompt_answer_ids = self.tokenizer.encode(prompt_answer)
|
||||
# answer_length = len(prompt_answer_ids) - prompt_pre_length
|
||||
labels_ids = copy.deepcopy(prompt_answer_ids) + [self.tokenizer.eos_token_id]
|
||||
labels_ids = torch.tensor(labels_ids, dtype=torch.int64) # [bos, prompt, answer, eos]
|
||||
labels_ids = torch.cat((audio_pseudo, labels_ids)) # [audio, bos, prompt, answer, eos]
|
||||
labels_ids[: audio_pseudo_length + prompt_pre_length] = -1 # [-1, -1, -1, answer, eos]
|
||||
# labels_ids[:prompt_pre_length] = -1 # [-1, -1, input, eos]
|
||||
label_mask = labels_ids.ge(0) # [false, false, false, true, true]
|
||||
labels_ids[~label_mask] = self.IGNORE_INDEX # [-100, -100, -100, answer, eos]
|
||||
|
||||
# audio_mask for input_ids
|
||||
audio_mask = [1] * audio_pseudo_length + [0] * input_ids_length
|
||||
audio_mask = torch.tensor(audio_mask, dtype=torch.float32)
|
||||
|
||||
ids = self.tokenizer.encode(target) # token ids is different from labels_ids
|
||||
text = torch.tensor(ids, dtype=torch.int64)
|
||||
text_lengths = torch.tensor([len(ids)], dtype=torch.int32)
|
||||
|
||||
return {
|
||||
"speech": speech,
|
||||
"speech_lengths": speech_lengths,
|
||||
"text": text,
|
||||
"text_lengths": text_lengths,
|
||||
"input_ids": input_ids,
|
||||
"attention_mask": attention_mask,
|
||||
"labels_ids": labels_ids,
|
||||
"label_mask": label_mask,
|
||||
"audio_mask": audio_mask,
|
||||
}
|
||||
|
||||
def collator(self, samples: list = None):
|
||||
"""Collator.
|
||||
|
||||
Args:
|
||||
samples: TODO.
|
||||
"""
|
||||
outputs = {}
|
||||
for sample in samples:
|
||||
for key in sample.keys():
|
||||
if key not in outputs:
|
||||
outputs[key] = []
|
||||
outputs[key].append(sample[key])
|
||||
|
||||
for key, data_list in outputs.items():
|
||||
if isinstance(data_list[0], torch.Tensor):
|
||||
if data_list[0].dtype == torch.int64 or data_list[0].dtype == torch.int32:
|
||||
|
||||
pad_value = self.int_pad_value
|
||||
else:
|
||||
pad_value = self.float_pad_value
|
||||
|
||||
outputs[key] = torch.nn.utils.rnn.pad_sequence(
|
||||
data_list, batch_first=True, padding_value=pad_value
|
||||
)
|
||||
return outputs
|
||||
@@ -0,0 +1,191 @@
|
||||
import torch
|
||||
import copy
|
||||
|
||||
from funasr.register import tables
|
||||
from funasr.utils.load_utils import extract_fbank, load_audio_text_image_video
|
||||
|
||||
|
||||
@tables.register("dataset_classes", "AudioLLMVicunaDataset")
|
||||
class AudioLLMVicunaDataset(torch.utils.data.Dataset):
|
||||
"""
|
||||
AudioLLMDataset
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path,
|
||||
index_ds: str = None,
|
||||
frontend=None,
|
||||
tokenizer=None,
|
||||
int_pad_value: int = -1,
|
||||
float_pad_value: float = 0.0,
|
||||
**kwargs
|
||||
):
|
||||
"""Initialize AudioLLMVicunaDataset.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
index_ds: TODO.
|
||||
frontend: Audio frontend for feature extraction.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
int_pad_value: TODO.
|
||||
float_pad_value: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
index_ds_class = tables.index_ds_classes.get(index_ds)
|
||||
self.index_ds = index_ds_class(path, **kwargs)
|
||||
preprocessor_speech = kwargs.get("preprocessor_speech", None)
|
||||
if preprocessor_speech:
|
||||
preprocessor_speech_class = tables.preprocessor_classes.get(preprocessor_speech)
|
||||
preprocessor_speech = preprocessor_speech_class(
|
||||
**kwargs.get("preprocessor_speech_conf", {})
|
||||
)
|
||||
self.preprocessor_speech = preprocessor_speech
|
||||
preprocessor_text = kwargs.get("preprocessor_text", None)
|
||||
if preprocessor_text:
|
||||
preprocessor_text_class = tables.preprocessor_classes.get(preprocessor_text)
|
||||
preprocessor_text = preprocessor_text_class(**kwargs.get("preprocessor_text_conf", {}))
|
||||
self.preprocessor_text = preprocessor_text
|
||||
|
||||
self.frontend = frontend
|
||||
self.fs = 16000 if frontend is None else frontend.fs
|
||||
self.data_type = "sound"
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.float_pad_value = float_pad_value
|
||||
self.prompt = kwargs.get("prompt", "Transcribe speech to text.")
|
||||
# self.prompt_pre = "USER: \nINSTRUCTION: {}\nINPUT: ".format(self.prompt) # "USER: \nINSTRUCTION: {}\nnINPUT: {}\nASSISTANT: "
|
||||
self.prompt_af = ""
|
||||
self.IGNORE_INDEX = kwargs.get("IGNORE_INDEX", -100)
|
||||
self.int_pad_value = self.IGNORE_INDEX
|
||||
self.audio_adaptor_downsample_rate = kwargs.get("audio_adaptor_downsample_rate", 5)
|
||||
self.audio_encoder_downsample_rate = kwargs.get("audio_encoder_downsample_rate", 2)
|
||||
self.prompt_template = "USER: {}\n ASSISTANT:"
|
||||
self.answer_template = "{}"
|
||||
|
||||
def get_source_len(self, index):
|
||||
"""Get source len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_source_len(item)
|
||||
|
||||
def get_target_len(self, index):
|
||||
"""Get target len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_target_len(item)
|
||||
|
||||
def __len__(self):
|
||||
"""Internal: len ."""
|
||||
return len(self.index_ds)
|
||||
|
||||
def __getitem__(self, index):
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
source = item["source"]
|
||||
data_src = load_audio_text_image_video(source, fs=self.fs)
|
||||
if self.preprocessor_speech:
|
||||
data_src = self.preprocessor_speech(data_src, fs=self.fs)
|
||||
speech, speech_lengths = extract_fbank(
|
||||
data_src, data_type=self.data_type, frontend=self.frontend, is_final=True
|
||||
) # speech: [b, T, d]
|
||||
speech = speech.squeeze(0)
|
||||
|
||||
audio_pseudo_length = (
|
||||
(speech.shape[0] + 1)
|
||||
// self.audio_adaptor_downsample_rate
|
||||
// self.audio_encoder_downsample_rate
|
||||
)
|
||||
audio_pseudo = torch.full((audio_pseudo_length,), -1) # placeholder
|
||||
|
||||
target = item["target"]
|
||||
if self.preprocessor_text:
|
||||
target = self.preprocessor_text(target)
|
||||
|
||||
self.prompt_pre = self.prompt_template.format(self.prompt)
|
||||
prompt_ids_pre = self.tokenizer.encode(self.prompt_pre) # [bos,prompt]
|
||||
prompt_pre_length = len(prompt_ids_pre)
|
||||
|
||||
# input
|
||||
input = self.answer_template.format(target.lower())
|
||||
prompt_input = "{}{}".format(self.prompt_pre, input)
|
||||
prompt_input_ids = self.tokenizer.encode(prompt_input) # [bos, prompt, input]
|
||||
# audio_length = len(prompt_input_ids) - prompt_pre_length
|
||||
input_ids = prompt_input_ids + [self.tokenizer.pad_token_id] # [bos, prompt, input, pad]
|
||||
input_ids_length = len(input_ids)
|
||||
input_ids = torch.tensor(input_ids, dtype=torch.int64) # [bos, prompt, input, pad]
|
||||
input_ids = torch.cat((audio_pseudo, input_ids)) # [audio, bos, prompt, input, pad]
|
||||
# input_ids[:audio_pseudo_length] = -1 # [-1, bos, prompt, input, pad]
|
||||
attention_mask = input_ids.ge(-1) # [true, true, true, true, true], length mask
|
||||
# input_ids[prompt_pre_length:] = -1 # [bos, prompt,-1,-1]
|
||||
# attention_mask = input_ids.ge(-1) # [true, true, true, true], length mask
|
||||
|
||||
# label
|
||||
answer = self.answer_template.format(target.lower())
|
||||
prompt_answer = "{}{}".format(self.prompt_pre, answer)
|
||||
prompt_answer_ids = self.tokenizer.encode(prompt_answer)
|
||||
# answer_length = len(prompt_answer_ids) - prompt_pre_length
|
||||
labels_ids = copy.deepcopy(prompt_answer_ids) + [self.tokenizer.eos_token_id]
|
||||
labels_ids = torch.tensor(labels_ids, dtype=torch.int64) # [bos, prompt, answer, eos]
|
||||
labels_ids = torch.cat((audio_pseudo, labels_ids)) # [audio, bos, prompt, answer, eos]
|
||||
labels_ids[: audio_pseudo_length + prompt_pre_length] = -1 # [-1, -1, -1, answer, eos]
|
||||
# labels_ids[:prompt_pre_length] = -1 # [-1, -1, input, eos]
|
||||
label_mask = labels_ids.ge(0) # [false, false, false, true, true]
|
||||
labels_ids[~label_mask] = self.IGNORE_INDEX # [-100, -100, -100, answer, eos]
|
||||
|
||||
# audio_mask for input_ids
|
||||
audio_mask = [1] * audio_pseudo_length + [0] * input_ids_length
|
||||
audio_mask = torch.tensor(audio_mask, dtype=torch.float32)
|
||||
|
||||
ids = self.tokenizer.encode(target) # token ids is different from labels_ids
|
||||
text = torch.tensor(ids, dtype=torch.int64)
|
||||
text_lengths = torch.tensor([len(ids)], dtype=torch.int32)
|
||||
|
||||
return {
|
||||
"speech": speech,
|
||||
"speech_lengths": speech_lengths,
|
||||
"text": text,
|
||||
"text_lengths": text_lengths,
|
||||
"input_ids": input_ids,
|
||||
"attention_mask": attention_mask,
|
||||
"labels_ids": labels_ids,
|
||||
"label_mask": label_mask,
|
||||
"audio_mask": audio_mask,
|
||||
}
|
||||
|
||||
def collator(self, samples: list = None):
|
||||
"""Collator.
|
||||
|
||||
Args:
|
||||
samples: TODO.
|
||||
"""
|
||||
outputs = {}
|
||||
for sample in samples:
|
||||
for key in sample.keys():
|
||||
if key not in outputs:
|
||||
outputs[key] = []
|
||||
outputs[key].append(sample[key])
|
||||
|
||||
for key, data_list in outputs.items():
|
||||
if isinstance(data_list[0], torch.Tensor):
|
||||
if data_list[0].dtype == torch.int64 or data_list[0].dtype == torch.int32:
|
||||
|
||||
pad_value = self.int_pad_value
|
||||
else:
|
||||
pad_value = self.float_pad_value
|
||||
|
||||
outputs[key] = torch.nn.utils.rnn.pad_sequence(
|
||||
data_list, batch_first=True, padding_value=pad_value
|
||||
)
|
||||
return outputs
|
||||
@@ -0,0 +1,545 @@
|
||||
import logging
|
||||
import re
|
||||
import torch
|
||||
import random
|
||||
import traceback
|
||||
from funasr.register import tables
|
||||
from funasr.utils.load_utils import extract_fbank, load_audio_text_image_video
|
||||
|
||||
|
||||
@tables.register("dataset_classes", "OpenAIDataset")
|
||||
class OpenAIDataset(torch.utils.data.Dataset):
|
||||
"""
|
||||
SenseVoiceDataset
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path,
|
||||
index_ds: str = None,
|
||||
frontend=None,
|
||||
tokenizer=None,
|
||||
int_pad_value: int = -1,
|
||||
float_pad_value: float = 0.0,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize OpenAIDataset.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
index_ds: TODO.
|
||||
frontend: Audio frontend for feature extraction.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
int_pad_value: TODO.
|
||||
float_pad_value: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
index_ds_class = tables.index_ds_classes.get(index_ds)
|
||||
self.index_ds = index_ds_class(path, **kwargs)
|
||||
preprocessor_speech = kwargs.get("preprocessor_speech", None)
|
||||
if preprocessor_speech:
|
||||
preprocessor_speech_class = tables.preprocessor_classes.get(preprocessor_speech)
|
||||
preprocessor_speech = preprocessor_speech_class(
|
||||
**kwargs.get("preprocessor_speech_conf")
|
||||
)
|
||||
self.preprocessor_speech = preprocessor_speech
|
||||
preprocessor_text = kwargs.get("preprocessor_text", None)
|
||||
if preprocessor_text:
|
||||
preprocessor_text_class = tables.preprocessor_classes.get(preprocessor_text)
|
||||
preprocessor_text = preprocessor_text_class(**kwargs.get("preprocessor_text_conf"))
|
||||
self.preprocessor_text = preprocessor_text
|
||||
|
||||
self.frontend = frontend
|
||||
self.fs = 16000 if frontend is None else frontend.fs
|
||||
self.data_type = "sound"
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.int_pad_value = int_pad_value
|
||||
self.float_pad_value = float_pad_value
|
||||
self.sos = kwargs.get("sos", "<|startoftranscript|>")
|
||||
self.eos = kwargs.get("eos", "<|endoftext|>")
|
||||
self.batch_size = kwargs.get("batch_size")
|
||||
self.batch_type = kwargs.get("batch_type")
|
||||
self.prompt_ids_len = 0
|
||||
self.retry = kwargs.get("retry", 100)
|
||||
|
||||
self.permute = False
|
||||
from funasr.frontends.whisper_frontend import WhisperFrontend
|
||||
|
||||
if isinstance(self.frontend, WhisperFrontend):
|
||||
self.permute = True
|
||||
|
||||
self.pattern = re.compile(r"(<\|startofspeech\|>.*?<\|endofspeech\|>)")
|
||||
# self.kwargs = kwargs
|
||||
self.max_token_length = kwargs.get("max_token_length", 1024)
|
||||
self.batch_size_scale_ratio_max = kwargs.get("batch_size_scale_ratio_max", 1.5)
|
||||
self.batch_size_token_max = kwargs.get("batch_size_token_max", 2500)
|
||||
self.audio_adaptor_downsample_rate = kwargs.get("audio_adaptor_downsample_rate", 2)
|
||||
self.audio_encoder_downsample_rate = kwargs.get("audio_encoder_downsample_rate", 4)
|
||||
|
||||
def get_source_len(self, index):
|
||||
"""Get source len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_source_len(item)
|
||||
|
||||
def get_target_len(self, index):
|
||||
"""Get target len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_target_len(item)
|
||||
|
||||
def __len__(self):
|
||||
"""Internal: len ."""
|
||||
return len(self.index_ds)
|
||||
|
||||
def __getitem__(self, index):
|
||||
# import pdb;
|
||||
# pdb.set_trace()
|
||||
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
output = None
|
||||
|
||||
for idx in range(self.retry):
|
||||
badcase_flag = False
|
||||
if idx == 0:
|
||||
index_cur = index
|
||||
else:
|
||||
index_cur = torch.randint(0, len(self.index_ds), ()).item()
|
||||
|
||||
item = self.index_ds[index_cur]
|
||||
|
||||
system = item["system"]
|
||||
user = item["user"]
|
||||
assistant = item["assistant"]
|
||||
|
||||
input_ids, labels, fbank, fbank_lens, fbank_mask, fbank_beg = [], [], [], [], [], []
|
||||
|
||||
for i, (system_prompt, user_prompt, target_out) in enumerate(
|
||||
zip(system, user, assistant)
|
||||
):
|
||||
|
||||
source_input = f"<|im_start|>system\n{system_prompt}<|im_end|>\n<|im_start|>user\n{user_prompt}<|im_end|>\n<|im_start|>assistant\n"
|
||||
|
||||
splits = self.pattern.split(source_input)
|
||||
source_ids = []
|
||||
fbank_mask_i = []
|
||||
fbank_beg_i = []
|
||||
fbank_lens_i = []
|
||||
for k, sub_str in enumerate(splits):
|
||||
if not sub_str.startswith("<|startofspeech|>"):
|
||||
sub_token = self.tokenizer.encode(sub_str)
|
||||
source_ids += sub_token
|
||||
fbank_mask_i += [0] * len(sub_token)
|
||||
else:
|
||||
sub_str = sub_str.replace("<|startofspeech|>", "").replace(
|
||||
"<|endofspeech|>", ""
|
||||
)
|
||||
if sub_str.startswith("!"):
|
||||
try:
|
||||
data_src = load_audio_text_image_video(sub_str[1:], fs=self.fs)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
f"Loading wav failed! {str(e)}, {traceback.format_exc()}"
|
||||
)
|
||||
badcase_flag = True
|
||||
continue
|
||||
speech, speech_lengths = extract_fbank(
|
||||
data_src,
|
||||
data_type=self.data_type,
|
||||
frontend=self.frontend,
|
||||
is_final=True,
|
||||
) # speech: [b, T, d]
|
||||
if self.permute:
|
||||
speech = speech.permute(0, 2, 1)
|
||||
# if speech_lengths > self.batch_size:
|
||||
# continue
|
||||
if self.audio_encoder_downsample_rate == 4:
|
||||
olens = 1 + (speech_lengths[0].item() - 3 + 2 * 1) // 2
|
||||
olens = 1 + (olens - 3 + 2 * 1) // 2
|
||||
elif self.audio_encoder_downsample_rate == 1:
|
||||
olens = speech_lengths[0].item()
|
||||
|
||||
sub_token_len = (olens - 1) // self.audio_adaptor_downsample_rate + 1
|
||||
sub_token = [0] * sub_token_len
|
||||
fbank_beg_i = [len(source_ids)]
|
||||
source_ids += sub_token
|
||||
fbank_mask_i += [1] * len(sub_token)
|
||||
|
||||
if badcase_flag:
|
||||
continue
|
||||
source_mask = [-100] * len(source_ids)
|
||||
target_out = f"{target_out}<|im_end|>"
|
||||
target_ids = self.tokenizer.encode(target_out)
|
||||
input_ids += source_ids + target_ids
|
||||
labels += source_mask + target_ids
|
||||
fbank_mask += fbank_mask_i
|
||||
fbank_beg.append(fbank_beg_i)
|
||||
|
||||
if len(input_ids) > self.max_token_length:
|
||||
logging.info(
|
||||
f"input_ids > max_token_length: {len(input_ids)}>{self.max_token_length}, {item}"
|
||||
)
|
||||
badcase_flag = True
|
||||
if badcase_flag:
|
||||
continue
|
||||
input_ids = torch.tensor(input_ids, dtype=torch.int64) # [: self.max_token_length]
|
||||
attention_mask = torch.tensor([1] * len(input_ids), dtype=torch.int32)
|
||||
labels = torch.tensor(labels, dtype=torch.int64) # [: self.max_token_length]
|
||||
|
||||
fbank = speech[0, :, :]
|
||||
fbank_lens = speech_lengths
|
||||
fbank_mask = torch.tensor(fbank_mask, dtype=torch.float32)
|
||||
fbank_beg = torch.tensor(fbank_beg, dtype=torch.int32)
|
||||
|
||||
output = {
|
||||
"speech": fbank,
|
||||
"speech_lengths": fbank_lens,
|
||||
"fbank_mask": fbank_mask,
|
||||
"fbank_beg": fbank_beg,
|
||||
"input_ids": input_ids,
|
||||
"attention_mask": attention_mask,
|
||||
"labels_ids": labels,
|
||||
}
|
||||
break
|
||||
|
||||
return output
|
||||
|
||||
def collator(self, samples: list = None):
|
||||
|
||||
"""Collator.
|
||||
|
||||
Args:
|
||||
samples: TODO.
|
||||
"""
|
||||
for idx in range(self.retry):
|
||||
badcase_flag = False
|
||||
|
||||
outputs = {}
|
||||
for sample in samples:
|
||||
if sample is None:
|
||||
continue
|
||||
for key in sample.keys():
|
||||
if key not in outputs:
|
||||
outputs[key] = []
|
||||
outputs[key].append(sample[key])
|
||||
|
||||
for key, data_list in outputs.items():
|
||||
if isinstance(data_list[0], torch.Tensor):
|
||||
if data_list[0].dtype == torch.int64 or data_list[0].dtype == torch.int32:
|
||||
|
||||
pad_value = self.int_pad_value
|
||||
else:
|
||||
pad_value = self.float_pad_value
|
||||
|
||||
outputs[key] = torch.nn.utils.rnn.pad_sequence(
|
||||
data_list, batch_first=True, padding_value=pad_value
|
||||
)
|
||||
|
||||
if self.batch_type != "example":
|
||||
b, t = outputs["input_ids"].shape
|
||||
if b > 1 and b * t > self.batch_size_token_max:
|
||||
logging.info(
|
||||
f"Warning, {idx}th, b*t: {b}*{t}={b * t} > batch_size_sample_max: {self.batch_size_token_max}, drop last data"
|
||||
)
|
||||
samples = samples[:-1]
|
||||
continue
|
||||
|
||||
break
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
@tables.register("dataset_classes", "OpenAIDatasetMultiTurn")
|
||||
class OpenAIDatasetMultiTurn(torch.utils.data.Dataset):
|
||||
"""
|
||||
SenseVoiceDataset
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path,
|
||||
index_ds: str = None,
|
||||
frontend=None,
|
||||
tokenizer=None,
|
||||
int_pad_value: int = -1,
|
||||
float_pad_value: float = 0.0,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize OpenAIDatasetMultiTurn.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
index_ds: TODO.
|
||||
frontend: Audio frontend for feature extraction.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
int_pad_value: TODO.
|
||||
float_pad_value: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
index_ds_class = tables.index_ds_classes.get(index_ds)
|
||||
self.index_ds = index_ds_class(path, **kwargs)
|
||||
preprocessor_speech = kwargs.get("preprocessor_speech", None)
|
||||
if preprocessor_speech:
|
||||
preprocessor_speech_class = tables.preprocessor_classes.get(preprocessor_speech)
|
||||
preprocessor_speech = preprocessor_speech_class(
|
||||
**kwargs.get("preprocessor_speech_conf")
|
||||
)
|
||||
self.preprocessor_speech = preprocessor_speech
|
||||
preprocessor_text = kwargs.get("preprocessor_text", None)
|
||||
if preprocessor_text:
|
||||
preprocessor_text_class = tables.preprocessor_classes.get(preprocessor_text)
|
||||
preprocessor_text = preprocessor_text_class(**kwargs.get("preprocessor_text_conf"))
|
||||
self.preprocessor_text = preprocessor_text
|
||||
|
||||
self.frontend = frontend
|
||||
self.fs = 16000 if frontend is None else frontend.fs
|
||||
self.data_type = "sound"
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.int_pad_value = int_pad_value
|
||||
self.float_pad_value = float_pad_value
|
||||
self.sos = kwargs.get("sos", "<|startoftranscript|>")
|
||||
self.eos = kwargs.get("eos", "<|endoftext|>")
|
||||
self.batch_size = kwargs.get("batch_size")
|
||||
self.batch_type = kwargs.get("batch_type")
|
||||
self.prompt_ids_len = 0
|
||||
self.retry = kwargs.get("retry", 100)
|
||||
|
||||
self.permute = False
|
||||
from funasr.frontends.whisper_frontend import WhisperFrontend
|
||||
|
||||
if isinstance(self.frontend, WhisperFrontend):
|
||||
self.permute = True
|
||||
|
||||
self.pattern = re.compile(r"(<\|startofspeech\|>.*?<\|endofspeech\|>)")
|
||||
# self.kwargs = kwargs
|
||||
self.max_token_length = kwargs.get("max_token_length", 1500)
|
||||
self.batch_size_scale_ratio_max = kwargs.get("batch_size_scale_ratio_max", 1.5)
|
||||
self.batch_size_token_max = kwargs.get("batch_size_token_max", 2500)
|
||||
self.multiturn_num_max = kwargs.get("multiturn_num_max", 5)
|
||||
self.max_source_length = kwargs.get("max_source_length", 3000)
|
||||
|
||||
def get_source_len(self, index):
|
||||
"""Get source len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_source_len(item)
|
||||
|
||||
def get_target_len(self, index):
|
||||
"""Get target len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_target_len(item)
|
||||
|
||||
def __len__(self):
|
||||
"""Internal: len ."""
|
||||
return len(self.index_ds)
|
||||
|
||||
def __getitem__(self, index):
|
||||
# import pdb
|
||||
#
|
||||
# pdb.set_trace()
|
||||
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
output = None
|
||||
|
||||
for idx in range(self.retry):
|
||||
badcase_flag = False
|
||||
if idx == 0:
|
||||
index_cur = index
|
||||
else:
|
||||
index_cur = torch.randint(0, len(self.index_ds), ()).item()
|
||||
|
||||
item = self.index_ds[index_cur]
|
||||
|
||||
system = item["system"]
|
||||
user = item["user"]
|
||||
assistant = item["assistant"]
|
||||
|
||||
input_ids, labels, fbank, fbank_lens, fbank_mask, fbank_beg, fake_token_len = (
|
||||
[],
|
||||
[],
|
||||
[],
|
||||
[],
|
||||
[],
|
||||
[],
|
||||
[],
|
||||
)
|
||||
|
||||
for i, (system_prompt, user_prompt, target_out) in enumerate(
|
||||
zip(system, user, assistant)
|
||||
):
|
||||
if i >= self.multiturn_num_max:
|
||||
break
|
||||
if len(input_ids) > self.max_token_length:
|
||||
logging.info(
|
||||
f"input_ids > max_token_length: {len(input_ids)}>{self.max_token_length}, {item}"
|
||||
)
|
||||
break
|
||||
|
||||
if i == 0:
|
||||
source_input = f"<|im_start|>system\n{system_prompt}<|im_end|>\n<|im_start|>user\n{user_prompt}<|im_end|>\n<|im_start|>assistant\n"
|
||||
else:
|
||||
source_input = (
|
||||
f"<|im_start|>user\n{user_prompt}<|im_end|>\n<|im_start|>assistant\n"
|
||||
)
|
||||
|
||||
splits = self.pattern.split(source_input)
|
||||
source_ids = []
|
||||
fbank_i = []
|
||||
fbank_mask_i = []
|
||||
fake_token_len_i = 0
|
||||
fbank_beg_i = -1
|
||||
fbank_lens_i = []
|
||||
for k, sub_str in enumerate(splits):
|
||||
if not sub_str.startswith("<|startofspeech|>"):
|
||||
sub_token = self.tokenizer.encode(sub_str)
|
||||
source_ids += sub_token
|
||||
fbank_mask_i += [0] * len(sub_token)
|
||||
else:
|
||||
sub_str = sub_str.replace("<|startofspeech|>", "").replace(
|
||||
"<|endofspeech|>", ""
|
||||
)
|
||||
if sub_str.startswith("!"):
|
||||
try:
|
||||
data_src = load_audio_text_image_video(sub_str[1:], fs=self.fs)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
f"Loading wav failed! {str(e)}, {traceback.format_exc()}"
|
||||
)
|
||||
badcase_flag = True
|
||||
continue
|
||||
speech, speech_lengths = extract_fbank(
|
||||
data_src,
|
||||
data_type=self.data_type,
|
||||
frontend=self.frontend,
|
||||
is_final=True,
|
||||
) # speech: [b, T, d]
|
||||
if speech_lengths > self.max_source_length:
|
||||
logging.info(
|
||||
f"speech_lengths > max_source_length: {speech_lengths}>{self.max_source_length}, {item}"
|
||||
)
|
||||
badcase_flag = True
|
||||
if self.permute:
|
||||
speech = speech.permute(0, 2, 1)
|
||||
# if speech_lengths > self.batch_size:
|
||||
# continue
|
||||
|
||||
olens = 1 + (speech_lengths[0].item() - 3 + 2 * 1) // 2
|
||||
olens = 1 + (olens - 3 + 2 * 1) // 2
|
||||
fake_token_len_i = (olens - 1) // 2 + 1
|
||||
fake_token = [0] * fake_token_len_i
|
||||
fbank_beg_i = len(source_ids)
|
||||
source_ids += fake_token
|
||||
fbank_mask_i += [1] * len(fake_token)
|
||||
|
||||
if badcase_flag:
|
||||
continue
|
||||
|
||||
fbank_beg += [fbank_beg_i + len(input_ids)]
|
||||
fake_token_len += [fake_token_len_i]
|
||||
source_mask = [-100] * len(source_ids)
|
||||
target_out = f"{target_out}<|im_end|>"
|
||||
target_ids = self.tokenizer.encode(target_out)
|
||||
input_ids += source_ids + target_ids
|
||||
labels += source_mask + target_ids
|
||||
fbank.append(speech[0, :, :])
|
||||
fbank_mask += fbank_mask_i
|
||||
fbank_lens.append(speech_lengths)
|
||||
|
||||
if badcase_flag:
|
||||
continue
|
||||
|
||||
input_ids = torch.tensor(input_ids, dtype=torch.int64) # [: self.max_token_length]
|
||||
attention_mask = torch.tensor([1] * len(input_ids), dtype=torch.int32)
|
||||
labels = torch.tensor(labels, dtype=torch.int64) # [: self.max_token_length]
|
||||
|
||||
# fbank = speech[0, :, :]
|
||||
# fbank_lens = torch.tensor(fbank_lens, dtype=torch.int32)
|
||||
fbank_mask = torch.tensor(fbank_mask, dtype=torch.float32)
|
||||
fbank_beg = torch.tensor(fbank_beg, dtype=torch.int32)
|
||||
fake_token_len = torch.tensor(fake_token_len, dtype=torch.int32)
|
||||
|
||||
output = {
|
||||
"speech": fbank,
|
||||
"speech_lengths": fbank_lens,
|
||||
"fbank_mask": fbank_mask,
|
||||
"fbank_beg": fbank_beg,
|
||||
"fake_token_len": fake_token_len,
|
||||
"input_ids": input_ids,
|
||||
"attention_mask": attention_mask,
|
||||
"labels_ids": labels,
|
||||
}
|
||||
break
|
||||
|
||||
return output
|
||||
|
||||
def collator(self, samples: list = None):
|
||||
|
||||
"""Collator.
|
||||
|
||||
Args:
|
||||
samples: TODO.
|
||||
"""
|
||||
for idx in range(self.retry):
|
||||
badcase_flag = False
|
||||
|
||||
outputs = {}
|
||||
for sample in samples:
|
||||
if sample is None:
|
||||
continue
|
||||
for key in sample.keys():
|
||||
if key not in outputs:
|
||||
outputs[key] = []
|
||||
if isinstance(sample[key], (list, tuple)):
|
||||
outputs[key].extend(sample[key])
|
||||
else:
|
||||
outputs[key].append(sample[key])
|
||||
|
||||
for key, data_list in outputs.items():
|
||||
if isinstance(data_list[0], torch.Tensor):
|
||||
if data_list[0].dtype == torch.int64 or data_list[0].dtype == torch.int32:
|
||||
|
||||
pad_value = self.int_pad_value
|
||||
else:
|
||||
pad_value = self.float_pad_value
|
||||
|
||||
outputs[key] = torch.nn.utils.rnn.pad_sequence(
|
||||
data_list, batch_first=True, padding_value=pad_value
|
||||
)
|
||||
|
||||
if self.batch_type != "example":
|
||||
b, t = outputs["input_ids"].shape
|
||||
if b > 1 and b * t > self.batch_size_token_max:
|
||||
logging.info(
|
||||
f"Warning, {idx}th, b*t: {b}*{t}={b * t} > batch_size_sample_max: {self.batch_size_token_max}, drop last data"
|
||||
)
|
||||
samples = samples[:-1]
|
||||
continue
|
||||
|
||||
break
|
||||
|
||||
return outputs
|
||||
@@ -0,0 +1,138 @@
|
||||
import os
|
||||
import json
|
||||
import torch
|
||||
import logging
|
||||
|
||||
import librosa
|
||||
import random
|
||||
import torch.distributed as dist
|
||||
|
||||
from funasr.register import tables
|
||||
|
||||
|
||||
@tables.register("index_ds_classes", "OpenAIIndexDSJsonl")
|
||||
class OpenAIIndexDSJsonl(torch.utils.data.Dataset): # torch.utils.data.Dataset
|
||||
|
||||
def __init__(self, path: str, **kwargs):
|
||||
"""Initialize OpenAIIndexDSJsonl.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
self.max_source_length = kwargs.get("max_source_length", 3000)
|
||||
self.min_source_length = kwargs.get("min_source_length", 0)
|
||||
self.max_target_length = kwargs.get("max_target_length", 2048)
|
||||
self.min_target_length = kwargs.get("min_target_length", 0)
|
||||
self.max_token_length = kwargs.get("max_token_length", 2200)
|
||||
|
||||
is_training = kwargs.get("is_training", True)
|
||||
if not (path.endswith(".jsonl") or path.endswith(".json")):
|
||||
# jsonl list file
|
||||
data_split_num = kwargs.get("data_split_num", 1)
|
||||
data_split_i = kwargs.get("data_split_i", 0)
|
||||
|
||||
if not is_training:
|
||||
data_split_num = 1
|
||||
data_split_i = 0
|
||||
with open(path, encoding="utf-8") as fin:
|
||||
file_list_all = fin.readlines()
|
||||
|
||||
num_per_slice = (len(file_list_all) - 1) // data_split_num + 1 # 16
|
||||
file_list = file_list_all[
|
||||
data_split_i * num_per_slice : (data_split_i + 1) * num_per_slice
|
||||
]
|
||||
logging.info(
|
||||
f"is_training: {is_training}, data_split_num: {data_split_num}, data_split_i: {data_split_i}, \nfile_list: {file_list}, \nfile_list_all: {file_list_all}"
|
||||
)
|
||||
|
||||
else:
|
||||
file_list = [path]
|
||||
|
||||
contents = []
|
||||
for file_json in file_list:
|
||||
with open(file_json.strip(), encoding="utf-8") as fin:
|
||||
for line in fin:
|
||||
data_dict = json.loads(line.strip())
|
||||
data = data_dict["messages"]
|
||||
speech_length = data_dict.get("speech_length", -1) // 8
|
||||
text_length = data_dict.get("text_length", 0)
|
||||
if speech_length > self.max_source_length:
|
||||
logging.info(
|
||||
"speech_length: {speech_length} > {self.max_source_length}, drop it"
|
||||
)
|
||||
continue
|
||||
if text_length > self.max_target_length:
|
||||
continue
|
||||
|
||||
self.max_target_length = kwargs.get("max_target_length", 2048)
|
||||
|
||||
system, user, assistant = [], [], []
|
||||
for i, item in enumerate(data):
|
||||
role = item["role"]
|
||||
content = item["content"]
|
||||
if role == "system":
|
||||
system.append(content)
|
||||
elif role == "user":
|
||||
user.append(content)
|
||||
elif role == "assistant":
|
||||
assistant.append(content)
|
||||
|
||||
system = system * len(user)
|
||||
|
||||
contents_i = {
|
||||
"system": system,
|
||||
"user": user,
|
||||
"assistant": assistant,
|
||||
"source_len": speech_length + text_length,
|
||||
}
|
||||
contents.append(contents_i)
|
||||
|
||||
self.contents = contents
|
||||
|
||||
logging.info("total_num of samplers: {}, {}".format(len(self.contents), path))
|
||||
|
||||
def __len__(self):
|
||||
"""Internal: len ."""
|
||||
return len(self.contents)
|
||||
|
||||
def __getitem__(self, index):
|
||||
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
data = self.contents[index]
|
||||
|
||||
return data
|
||||
|
||||
def get_source_len(self, data_dict):
|
||||
"""Get source len.
|
||||
|
||||
Args:
|
||||
data_dict: TODO.
|
||||
"""
|
||||
source_len = data_dict.get("source_len", -1)
|
||||
if source_len < 0:
|
||||
source_len = len(data_dict["system"]) + len(data_dict["user"])
|
||||
return source_len
|
||||
|
||||
def get_target_len(self, data_dict):
|
||||
|
||||
"""Get target len.
|
||||
|
||||
Args:
|
||||
data_dict: TODO.
|
||||
"""
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
index_ds = OpenAIIndexDSJsonl(
|
||||
path="/Users/zhifu/funasr1.0/test_local/data_tmp/tmp_wav_10.jsonl"
|
||||
)
|
||||
print(index_ds.contents)
|
||||
pass
|
||||
@@ -0,0 +1,506 @@
|
||||
import logging
|
||||
|
||||
import re
|
||||
import torch
|
||||
import random
|
||||
import traceback
|
||||
from funasr.register import tables
|
||||
from funasr.utils.load_utils import extract_fbank, load_audio_text_image_video
|
||||
|
||||
|
||||
@tables.register("dataset_classes", "SenseVoiceDataset")
|
||||
class SenseVoiceDataset(torch.utils.data.Dataset):
|
||||
"""
|
||||
SenseVoiceDataset
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path,
|
||||
index_ds: str = None,
|
||||
frontend=None,
|
||||
tokenizer=None,
|
||||
int_pad_value: int = -1,
|
||||
float_pad_value: float = 0.0,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize SenseVoiceDataset.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
index_ds: TODO.
|
||||
frontend: Audio frontend for feature extraction.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
int_pad_value: TODO.
|
||||
float_pad_value: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
index_ds_class = tables.index_ds_classes.get(index_ds)
|
||||
self.index_ds = index_ds_class(path, **kwargs)
|
||||
preprocessor_speech = kwargs.get("preprocessor_speech", None)
|
||||
if preprocessor_speech:
|
||||
preprocessor_speech_class = tables.preprocessor_classes.get(preprocessor_speech)
|
||||
preprocessor_speech = preprocessor_speech_class(
|
||||
**kwargs.get("preprocessor_speech_conf")
|
||||
)
|
||||
self.preprocessor_speech = preprocessor_speech
|
||||
preprocessor_text = kwargs.get("preprocessor_text", None)
|
||||
if preprocessor_text:
|
||||
preprocessor_text_class = tables.preprocessor_classes.get(preprocessor_text)
|
||||
preprocessor_text = preprocessor_text_class(**kwargs.get("preprocessor_text_conf"))
|
||||
self.preprocessor_text = preprocessor_text
|
||||
|
||||
self.frontend = frontend
|
||||
self.fs = 16000 if frontend is None else frontend.fs
|
||||
self.data_type = "sound"
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.int_pad_value = int_pad_value
|
||||
self.float_pad_value = float_pad_value
|
||||
self.sos = kwargs.get("sos", "<|startoftranscript|>")
|
||||
self.eos = kwargs.get("eos", "<|endoftext|>")
|
||||
self.batch_size = kwargs.get("batch_size")
|
||||
self.batch_type = kwargs.get("batch_type")
|
||||
self.prompt_ids_len = 0
|
||||
self.retry = kwargs.get("retry", 5)
|
||||
|
||||
self.permute = False
|
||||
from funasr.frontends.whisper_frontend import WhisperFrontend
|
||||
|
||||
if isinstance(self.frontend, WhisperFrontend):
|
||||
self.permute = True
|
||||
|
||||
def get_source_len(self, index):
|
||||
"""Get source len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_source_len(item)
|
||||
|
||||
def get_target_len(self, index):
|
||||
"""Get target len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_target_len(item)
|
||||
|
||||
def __len__(self):
|
||||
"""Internal: len ."""
|
||||
return len(self.index_ds)
|
||||
|
||||
def __getitem__(self, index):
|
||||
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
output = None
|
||||
for idx in range(self.retry):
|
||||
if idx == 0:
|
||||
index_cur = index
|
||||
else:
|
||||
index_cur = torch.randint(0, len(self.index_ds), ()).item()
|
||||
|
||||
item = self.index_ds[index_cur]
|
||||
|
||||
source = item["source"]
|
||||
try:
|
||||
data_src = load_audio_text_image_video(source, fs=self.fs)
|
||||
except Exception as e:
|
||||
logging.error(f"Loading wav failed! {str(e)}, {traceback.format_exc()}")
|
||||
continue
|
||||
|
||||
if self.preprocessor_speech:
|
||||
data_src = self.preprocessor_speech(data_src, fs=self.fs)
|
||||
speech, speech_lengths = extract_fbank(
|
||||
data_src, data_type=self.data_type, frontend=self.frontend, is_final=True
|
||||
) # speech: [b, T, d]
|
||||
|
||||
if speech_lengths > self.batch_size:
|
||||
continue
|
||||
if self.permute:
|
||||
speech = speech.permute(0, 2, 1)
|
||||
target = item["target"]
|
||||
if self.preprocessor_text:
|
||||
target = self.preprocessor_text(target)
|
||||
|
||||
task = item.get("prompt", "<|ASR|>")
|
||||
text_language = item.get("text_language", "<|zh|>")
|
||||
|
||||
if isinstance(self.sos, str):
|
||||
prompt = f"{self.sos}{task}{text_language}"
|
||||
prompt_ids = self.tokenizer.encode(prompt, allowed_special="all")
|
||||
else:
|
||||
prompt = f"{task}{text_language}"
|
||||
prompt_ids = self.tokenizer.encode(prompt, allowed_special="all")
|
||||
prompt_ids = [self.sos] + prompt_ids
|
||||
|
||||
prompt_ids_len = len(prompt_ids) - 1 # [sos, task]
|
||||
self.prompt_ids_len = prompt_ids_len
|
||||
|
||||
target_ids = self.tokenizer.encode(target, allowed_special="all")
|
||||
target_ids_len = len(target_ids) + 1 # [lid, text]
|
||||
if target_ids_len > 200:
|
||||
continue
|
||||
|
||||
if isinstance(self.eos, str):
|
||||
eos = self.tokenizer.encode(self.eos, allowed_special="all") # [eos]
|
||||
else:
|
||||
eos = [self.eos]
|
||||
|
||||
ids = prompt_ids + target_ids + eos # [sos, task, lid, text, eos]
|
||||
ids_lengths = len(ids)
|
||||
|
||||
text = torch.tensor(ids, dtype=torch.int64)
|
||||
text_lengths = torch.tensor([ids_lengths], dtype=torch.int32)
|
||||
|
||||
target_mask = (
|
||||
[0] * (prompt_ids_len) + [1] * (target_ids_len) + [1]
|
||||
) # [sos, task, lid, text, eos]: [0, 0, 1, 1, 1]
|
||||
target_mask_lengths = len(target_mask)
|
||||
target_mask = torch.tensor(target_mask, dtype=torch.float32)
|
||||
target_mask_lengths = torch.tensor([target_mask_lengths], dtype=torch.int32)
|
||||
|
||||
output = {
|
||||
"speech": speech[0, :, :],
|
||||
"speech_lengths": speech_lengths,
|
||||
"text": text,
|
||||
"text_lengths": text_lengths,
|
||||
"target_mask": target_mask,
|
||||
"target_mask_lengths": target_mask_lengths,
|
||||
}
|
||||
break
|
||||
|
||||
return output
|
||||
|
||||
def collator(self, samples: list = None):
|
||||
"""Collator.
|
||||
|
||||
Args:
|
||||
samples: TODO.
|
||||
"""
|
||||
outputs = {}
|
||||
for sample in samples:
|
||||
if sample is None:
|
||||
continue
|
||||
for key in sample.keys():
|
||||
if key not in outputs:
|
||||
outputs[key] = []
|
||||
outputs[key].append(sample[key])
|
||||
|
||||
if len(outputs) < 1:
|
||||
logging.error(f"ERROR: data is empty!")
|
||||
outputs = {
|
||||
"speech": torch.rand((10, 128), dtype=torch.float32)[None, :, :],
|
||||
"speech_lengths": torch.tensor(
|
||||
[
|
||||
10,
|
||||
],
|
||||
dtype=torch.int32,
|
||||
)[:, None],
|
||||
"text": torch.tensor(
|
||||
[
|
||||
58836,
|
||||
],
|
||||
dtype=torch.int32,
|
||||
)[None, :],
|
||||
"text_lengths": torch.tensor(
|
||||
[
|
||||
1,
|
||||
],
|
||||
dtype=torch.int32,
|
||||
)[:, None],
|
||||
"target_mask": torch.tensor([[0] * (self.prompt_ids_len) + [1] * (1) + [1]])[
|
||||
None, :
|
||||
],
|
||||
}
|
||||
return outputs
|
||||
|
||||
for key, data_list in outputs.items():
|
||||
if isinstance(data_list[0], torch.Tensor):
|
||||
if data_list[0].dtype == torch.int64 or data_list[0].dtype == torch.int32:
|
||||
|
||||
pad_value = self.int_pad_value
|
||||
else:
|
||||
pad_value = self.float_pad_value
|
||||
|
||||
outputs[key] = torch.nn.utils.rnn.pad_sequence(
|
||||
data_list, batch_first=True, padding_value=pad_value
|
||||
)
|
||||
|
||||
if self.batch_type != "example":
|
||||
for i in range(10):
|
||||
outputs = self._filter_badcase(outputs, i=i)
|
||||
|
||||
return outputs
|
||||
|
||||
def _filter_badcase(self, outputs, i=0):
|
||||
"""Internal: filter badcase.
|
||||
|
||||
Args:
|
||||
outputs: TODO.
|
||||
i: TODO.
|
||||
"""
|
||||
b, t, _ = outputs["speech"].shape
|
||||
|
||||
if b * t > self.batch_size * 1.25:
|
||||
beg = torch.randint(0, 2, ()).item()
|
||||
if b < 2:
|
||||
beg = 0
|
||||
logging.info(
|
||||
f"Warning, b * t: {b * t} > {self.batch_size}, drop half data {i}th, beg:{beg}"
|
||||
)
|
||||
for key, data_list in outputs.items():
|
||||
outputs[key] = outputs[key][beg : beg + b : 2]
|
||||
|
||||
speech_lengths_max = outputs["speech_lengths"].max().item()
|
||||
outputs["speech"] = outputs["speech"][:, :speech_lengths_max, :]
|
||||
text_lengths_max = outputs["text_lengths"].max().item()
|
||||
outputs["text"] = outputs["text"][:, :text_lengths_max]
|
||||
target_mask_lengths_max = outputs["target_mask_lengths"].max().item()
|
||||
outputs["target_mask"] = outputs["target_mask"][:, :target_mask_lengths_max]
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
@tables.register("dataset_classes", "SenseVoiceCTCDataset")
|
||||
class SenseVoiceCTCDataset(torch.utils.data.Dataset):
|
||||
"""
|
||||
SenseVoiceCTCDataset
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path,
|
||||
index_ds: str = None,
|
||||
frontend=None,
|
||||
tokenizer=None,
|
||||
int_pad_value: int = -1,
|
||||
float_pad_value: float = 0.0,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize SenseVoiceCTCDataset.
|
||||
|
||||
Args:
|
||||
path: TODO.
|
||||
index_ds: TODO.
|
||||
frontend: Audio frontend for feature extraction.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
int_pad_value: TODO.
|
||||
float_pad_value: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
index_ds_class = tables.index_ds_classes.get(index_ds)
|
||||
self.index_ds = index_ds_class(path, **kwargs)
|
||||
preprocessor_speech = kwargs.get("preprocessor_speech", None)
|
||||
if preprocessor_speech:
|
||||
preprocessor_speech_class = tables.preprocessor_classes.get(preprocessor_speech)
|
||||
preprocessor_speech = preprocessor_speech_class(
|
||||
**kwargs.get("preprocessor_speech_conf")
|
||||
)
|
||||
self.preprocessor_speech = preprocessor_speech
|
||||
preprocessor_text = kwargs.get("preprocessor_text", None)
|
||||
if preprocessor_text:
|
||||
preprocessor_text_class = tables.preprocessor_classes.get(preprocessor_text)
|
||||
preprocessor_text = preprocessor_text_class(**kwargs.get("preprocessor_text_conf"))
|
||||
self.preprocessor_text = preprocessor_text
|
||||
|
||||
self.frontend = frontend
|
||||
self.fs = 16000 if frontend is None else frontend.fs
|
||||
self.data_type = "sound"
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.int_pad_value = int_pad_value
|
||||
self.float_pad_value = float_pad_value
|
||||
self.sos = kwargs.get("sos", "<|startoftranscript|>")
|
||||
self.eos = kwargs.get("eos", "<|endoftext|>")
|
||||
self.batch_size = kwargs.get("batch_size")
|
||||
self.batch_type = kwargs.get("batch_type")
|
||||
self.prompt_ids_len = 0
|
||||
self.retry = kwargs.get("retry", 5)
|
||||
|
||||
self.permute = False
|
||||
from funasr.frontends.whisper_frontend import WhisperFrontend
|
||||
|
||||
if isinstance(self.frontend, WhisperFrontend):
|
||||
self.permute = True
|
||||
|
||||
def get_source_len(self, index):
|
||||
"""Get source len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_source_len(item)
|
||||
|
||||
def get_target_len(self, index):
|
||||
"""Get target len.
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
item = self.index_ds[index]
|
||||
return self.index_ds.get_target_len(item)
|
||||
|
||||
def __len__(self):
|
||||
"""Internal: len ."""
|
||||
return len(self.index_ds)
|
||||
|
||||
def __getitem__(self, index):
|
||||
|
||||
"""Internal: getitem .
|
||||
|
||||
Args:
|
||||
index: TODO.
|
||||
"""
|
||||
output = None
|
||||
for idx in range(self.retry):
|
||||
if idx == 0:
|
||||
index_cur = index
|
||||
else:
|
||||
index_cur = torch.randint(0, len(self.index_ds), ()).item()
|
||||
|
||||
item = self.index_ds[index_cur]
|
||||
|
||||
source = item["source"]
|
||||
try:
|
||||
data_src = load_audio_text_image_video(source, fs=self.fs)
|
||||
except Exception as e:
|
||||
logging.error(f"Loading wav failed! {str(e)}, {traceback.format_exc()}")
|
||||
continue
|
||||
|
||||
if self.preprocessor_speech:
|
||||
data_src = self.preprocessor_speech(data_src, fs=self.fs)
|
||||
speech, speech_lengths = extract_fbank(
|
||||
data_src, data_type=self.data_type, frontend=self.frontend, is_final=True
|
||||
) # speech: [b, T, d]
|
||||
|
||||
if speech_lengths > self.batch_size:
|
||||
continue
|
||||
if self.permute:
|
||||
speech = speech.permute(0, 2, 1)
|
||||
asr_target = item["target"]
|
||||
if self.preprocessor_text:
|
||||
asr_target = self.preprocessor_text(asr_target)
|
||||
emo_target = item.get("emo_target", "<|NEUTRAL|>")
|
||||
event_target = item.get("event_target", "<|Speech|>")
|
||||
text_language = item.get("text_language", "<|zh|>")
|
||||
punc_itn_bottom = item.get("with_or_wo_itn", "<|woitn|>")
|
||||
|
||||
target_ids = self.tokenizer.encode(asr_target, allowed_special="all")
|
||||
target_ids_len = len(target_ids) # [text]
|
||||
if target_ids_len > 200:
|
||||
continue
|
||||
|
||||
lid_ids = self.tokenizer.encode(text_language, allowed_special="all")
|
||||
emo_ids = self.tokenizer.encode(emo_target, allowed_special="all")
|
||||
event_ids = self.tokenizer.encode(event_target, allowed_special="all")
|
||||
punc_itn_bottom_ids = self.tokenizer.encode(punc_itn_bottom, allowed_special="all")
|
||||
|
||||
ids = lid_ids + emo_ids + event_ids + punc_itn_bottom_ids + target_ids # [lid, emo, lid, itn, text]
|
||||
ids_lengths = len(ids)
|
||||
|
||||
text = torch.tensor(ids, dtype=torch.int64)
|
||||
text_lengths = torch.tensor([ids_lengths], dtype=torch.int32)
|
||||
|
||||
output = {
|
||||
"speech": speech[0, :, :],
|
||||
"speech_lengths": speech_lengths,
|
||||
"text": text,
|
||||
"text_lengths": text_lengths,
|
||||
}
|
||||
break
|
||||
|
||||
return output
|
||||
|
||||
def collator(self, samples: list = None):
|
||||
"""Collator.
|
||||
|
||||
Args:
|
||||
samples: TODO.
|
||||
"""
|
||||
outputs = {}
|
||||
for sample in samples:
|
||||
if sample is None:
|
||||
continue
|
||||
for key in sample.keys():
|
||||
if key not in outputs:
|
||||
outputs[key] = []
|
||||
outputs[key].append(sample[key])
|
||||
|
||||
if len(outputs) < 1:
|
||||
logging.error(f"ERROR: data is empty!")
|
||||
outputs = {
|
||||
"speech": torch.rand((10, 128), dtype=torch.float32)[None, :, :],
|
||||
"speech_lengths": torch.tensor(
|
||||
[
|
||||
10,
|
||||
],
|
||||
dtype=torch.int32,
|
||||
)[:, None],
|
||||
"text": torch.tensor(
|
||||
[
|
||||
58836,
|
||||
],
|
||||
dtype=torch.int32,
|
||||
)[None, :],
|
||||
"text_lengths": torch.tensor(
|
||||
[
|
||||
1,
|
||||
],
|
||||
dtype=torch.int32,
|
||||
)[:, None],
|
||||
}
|
||||
return outputs
|
||||
|
||||
for key, data_list in outputs.items():
|
||||
if isinstance(data_list[0], torch.Tensor):
|
||||
if data_list[0].dtype == torch.int64 or data_list[0].dtype == torch.int32:
|
||||
|
||||
pad_value = self.int_pad_value
|
||||
else:
|
||||
pad_value = self.float_pad_value
|
||||
|
||||
outputs[key] = torch.nn.utils.rnn.pad_sequence(
|
||||
data_list, batch_first=True, padding_value=pad_value
|
||||
)
|
||||
|
||||
if self.batch_type != "example":
|
||||
for i in range(10):
|
||||
outputs = self._filter_badcase(outputs, i=i)
|
||||
|
||||
return outputs
|
||||
|
||||
def _filter_badcase(self, outputs, i=0):
|
||||
"""Internal: filter badcase.
|
||||
|
||||
Args:
|
||||
outputs: TODO.
|
||||
i: TODO.
|
||||
"""
|
||||
b, t, _ = outputs["speech"].shape
|
||||
|
||||
if b * t > self.batch_size * 1.25:
|
||||
beg = torch.randint(0, 2, ()).item()
|
||||
if b < 2:
|
||||
beg = 0
|
||||
logging.info(
|
||||
f"Warning, b * t: {b * t} > {self.batch_size}, drop half data {i}th, beg:{beg}"
|
||||
)
|
||||
for key, data_list in outputs.items():
|
||||
outputs[key] = outputs[key][beg : beg + b : 2]
|
||||
|
||||
speech_lengths_max = outputs["speech_lengths"].max().item()
|
||||
outputs["speech"] = outputs["speech"][:, :speech_lengths_max, :]
|
||||
text_lengths_max = outputs["text_lengths"].max().item()
|
||||
outputs["text"] = outputs["text"][:, :text_lengths_max]
|
||||
|
||||
return outputs
|
||||
Reference in New Issue
Block a user