| 897 | |
| 898 | class DynamicMixingPreprocessor(AbsPreprocessor): |
| 899 | def __init__( |
| 900 | self, |
| 901 | train: bool, |
| 902 | source_scp: Optional[str] = None, |
| 903 | ref_num: int = 2, |
| 904 | dynamic_mixing_gain_db: float = 0.0, |
| 905 | speech_name: str = "speech_mix", |
| 906 | speech_ref_name_prefix: str = "speech_ref", |
| 907 | mixture_source_name: Optional[str] = None, |
| 908 | utt2spk: Optional[str] = None, |
| 909 | categories: Optional[List] = None, |
| 910 | ): |
| 911 | super().__init__(train) |
| 912 | self.source_scp = source_scp |
| 913 | self.ref_num = ref_num |
| 914 | self.dynamic_mixing_gain_db = dynamic_mixing_gain_db |
| 915 | self.speech_name = speech_name |
| 916 | self.speech_ref_name_prefix = speech_ref_name_prefix |
| 917 | # mixture_source_name: the key to select source utterances from dataloader |
| 918 | if mixture_source_name is None: |
| 919 | self.mixture_source_name = f"{speech_ref_name_prefix}1" |
| 920 | else: |
| 921 | self.mixture_source_name = mixture_source_name |
| 922 | |
| 923 | self.sources = {} |
| 924 | assert ( |
| 925 | source_scp is not None |
| 926 | ), f"Please pass `source_scp` to {type(self).__name__}" |
| 927 | with open(source_scp, "r", encoding="utf-8") as f: |
| 928 | for line in f: |
| 929 | sps = line.strip().split(None, 1) |
| 930 | assert len(sps) == 2 |
| 931 | self.sources[sps[0]] = sps[1] |
| 932 | |
| 933 | self.utt2spk = {} |
| 934 | if utt2spk is None: |
| 935 | # if utt2spk is not provided, create a dummy utt2spk with uid. |
| 936 | for key in self.sources.keys(): |
| 937 | self.utt2spk[key] = key |
| 938 | else: |
| 939 | with open(utt2spk, "r", encoding="utf-8") as f: |
| 940 | for line in f: |
| 941 | sps = line.strip().split(None, 1) |
| 942 | assert len(sps) == 2 |
| 943 | self.utt2spk[sps[0]] = sps[1] |
| 944 | |
| 945 | for key in self.sources.keys(): |
| 946 | assert key in self.utt2spk |
| 947 | |
| 948 | self.source_keys = list(self.sources.keys()) |
| 949 | |
| 950 | # Map each category into a unique integer |
| 951 | self.categories = {} |
| 952 | if categories: |
| 953 | count = 0 |
| 954 | for c in categories: |
| 955 | if c not in self.categories: |
| 956 | self.categories[c] = count |