Initial commit: FunASR Speech Recognition Toolkit
Update API Documentation / build-api-docs (push) Has been cancelled

Add complete FunASR codebase including models, runtime, and documentation.
This commit is contained in:
freedakgmail
2026-07-09 22:38:58 +08:00
commit 6116b1f3c6
3683 changed files with 990984 additions and 0 deletions
View File
+330
View File
@@ -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
+173
View File
@@ -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
+592
View File
@@ -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
+148
View File
@@ -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()
+141
View File
@@ -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()
+168
View File
@@ -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
+165
View File
@@ -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
+530
View File
@@ -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
+545
View File
@@ -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
+138
View File
@@ -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