bezzam HF Staff commited on
Commit
568d596
·
verified ·
1 Parent(s): 4e12205

Delete nemo_asr

Browse files
nemo_asr/run_eval.py DELETED
@@ -1,241 +0,0 @@
1
- import argparse
2
-
3
- import io
4
- import os
5
- import torch
6
- import evaluate
7
- import soundfile
8
-
9
- from tqdm import tqdm
10
- from normalizer import data_utils
11
- import numpy as np
12
-
13
- from nemo.collections.asr.models import ASRModel
14
- import time
15
-
16
-
17
- wer_metric = evaluate.load("wer")
18
-
19
-
20
- def main(args):
21
-
22
- data_cache_root = args.data_cache_root if args.data_cache_root is not None else os.getcwd()
23
- DATA_CACHE_DIR = os.path.join(data_cache_root, "audio_cache")
24
- DATASET_NAME = args.dataset
25
- SPLIT_NAME = args.split
26
-
27
- CACHE_DIR = os.path.join(DATA_CACHE_DIR, DATASET_NAME, SPLIT_NAME)
28
- if not os.path.exists(CACHE_DIR):
29
- os.makedirs(CACHE_DIR)
30
-
31
- if args.device >= 0:
32
- device = torch.device(f"cuda:{args.device}")
33
- compute_dtype=torch.bfloat16
34
- else:
35
- device = torch.device("cpu")
36
- compute_dtype=torch.float32
37
-
38
- if args.model_id.endswith(".nemo"):
39
- asr_model = ASRModel.restore_from(args.model_id, map_location=device)
40
- else:
41
- asr_model = ASRModel.from_pretrained(args.model_id, map_location=device) # type: ASRModel
42
-
43
- asr_model.to(compute_dtype)
44
- asr_model.eval()
45
- print(f"Model size: {sum(p.numel() for p in asr_model.parameters()) / 1e9:.2f}B parameters")
46
-
47
- dataset = data_utils.load_data(args)
48
-
49
- if args.max_eval_samples is not None and args.max_eval_samples > 0:
50
- print(f"Subsampling dataset to first {args.max_eval_samples} samples !")
51
- dataset = dataset.take(args.max_eval_samples)
52
-
53
- # Prepare data FIRST - this casts audio to proper format with "array" and "sampling_rate" keys
54
- dataset = data_utils.prepare_data(dataset)
55
-
56
- def download_audio_files(batch):
57
-
58
- # download audio files and write the paths, transcriptions and durations to a manifest file
59
- audio_paths = []
60
- original_audio_paths = []
61
- durations = []
62
- file_names = batch.get("file_name", [None] * len(batch["audio"]))
63
-
64
- # Use 'id' column if available, otherwise generate sequential IDs
65
- if "id" in batch:
66
- ids = batch["id"]
67
- else:
68
- # Generate IDs based on index
69
- start_idx = len([f for f in os.listdir(CACHE_DIR) if f.endswith('.wav')]) if os.path.exists(CACHE_DIR) else 0
70
- ids = [f"sample_{start_idx + i}" for i in range(len(batch["audio"]))]
71
-
72
- for id, file_name, audio_sample in zip(ids, file_names, batch["audio"]):
73
-
74
- # first step added here to make ID and wav filenames unique
75
- # several datasets like earnings22 have a hierarchical structure
76
- # for eg. earnings22/test/4432298/281.wav, earnings22/test/4450488/281.wav
77
- # lhotse uses the filename (281.wav) here as unique ID to create and name cuts
78
- # ref: https://github.com/lhotse-speech/lhotse/blob/master/lhotse/dataset/collation.py#L186
79
- original_id = id # preserve before sanitization for use as audio_filepath
80
- id = id.replace('/', '_').removesuffix('.wav')
81
-
82
- audio_path = os.path.join(CACHE_DIR, f"{id}.wav")
83
- audio_array = np.float32(audio_sample["array"])
84
- sample_rate = audio_sample["sampling_rate"]
85
-
86
- if not os.path.exists(audio_path):
87
- os.makedirs(os.path.dirname(audio_path), exist_ok=True)
88
- soundfile.write(audio_path, audio_array, sample_rate)
89
-
90
- audio_paths.append(audio_path)
91
- # Prefer the original file_name from the dataset; fall back to the
92
- # sample id (before path-sanitization) so audio_filepath in the
93
- # JSONL is always a meaningful identifier rather than "sample_N".
94
- if file_name is not None:
95
- original_audio_paths.append(os.path.basename(str(file_name)))
96
- else:
97
- original_audio_paths.append(original_id)
98
- durations.append(len(audio_array) / sample_rate)
99
-
100
-
101
- batch["references"] = batch["norm_text"]
102
- batch["audio_filepaths"] = audio_paths
103
- batch["original_audio_filepaths"] = original_audio_paths
104
- batch["durations"] = durations
105
-
106
- return batch
107
-
108
- if asr_model.cfg.decoding.strategy != "beam":
109
- asr_model.cfg.decoding.strategy = "greedy_batch"
110
- asr_model.change_decoding_strategy(asr_model.cfg.decoding)
111
-
112
- # prepraing the offline dataset
113
- dataset = dataset.map(download_audio_files, batch_size=args.batch_size, batched=True, remove_columns=["audio"])
114
-
115
- # Write manifest from daraset batch using json and keys audio_filepath, duration, text
116
-
117
- all_data = {
118
- "audio_filepaths": [],
119
- "original_audio_filepaths": [],
120
- "durations": [],
121
- "references": [],
122
- }
123
-
124
- data_itr = iter(dataset)
125
- for data in tqdm(data_itr, desc="Downloading Samples"):
126
- for key in all_data:
127
- all_data[key].append(data[key])
128
-
129
- # Sort audio_filepaths and references based on durations values
130
- sorted_indices = sorted(range(len(all_data["durations"])), key=lambda k: all_data["durations"][k], reverse=True)
131
- all_data["audio_filepaths"] = [all_data["audio_filepaths"][i] for i in sorted_indices]
132
- all_data["original_audio_filepaths"] = [all_data["original_audio_filepaths"][i] for i in sorted_indices]
133
- all_data["references"] = [all_data["references"][i] for i in sorted_indices]
134
- all_data["durations"] = [all_data["durations"][i] for i in sorted_indices]
135
-
136
-
137
- total_time = 0
138
- for _ in range(2): # warmup once and calculate rtf
139
- if _ == 0:
140
- audio_files = all_data["audio_filepaths"][:args.batch_size * 4] # warmup with 4 batches
141
- else:
142
- audio_files = all_data["audio_filepaths"]
143
- start_time = time.time()
144
- with torch.inference_mode(), torch.no_grad():
145
-
146
- if 'canary' in args.model_id and 'v2' not in args.model_id:
147
- pnc = 'nopnc'
148
- else:
149
- pnc = 'pnc'
150
-
151
- if 'canary' in args.model_id:
152
- transcriptions = asr_model.transcribe(audio_files, batch_size=args.batch_size, verbose=False, pnc=pnc, num_workers=1)
153
- else:
154
- transcriptions = asr_model.transcribe(audio_files, batch_size=args.batch_size, verbose=False, num_workers=1)
155
- end_time = time.time()
156
- if _ == 1:
157
- total_time += end_time - start_time
158
- total_time = total_time
159
-
160
- # normalize transcriptions with English normalizer
161
- if isinstance(transcriptions, tuple) and len(transcriptions) == 2:
162
- transcriptions = transcriptions[0]
163
- predictions = [data_utils.normalizer(pred.text) for pred in transcriptions]
164
-
165
- avg_time = total_time / len(all_data["audio_filepaths"])
166
-
167
- # Write manifest results (WER and RTFX)
168
- manifest_path = data_utils.write_manifest(
169
- all_data["references"],
170
- predictions,
171
- args.model_id,
172
- args.dataset_path,
173
- args.dataset,
174
- args.split,
175
- audio_length=all_data["durations"],
176
- transcription_time=[avg_time] * len(all_data["audio_filepaths"]),
177
- audio_filepaths=all_data["original_audio_filepaths"],
178
- )
179
-
180
- print("Results saved at path:", os.path.abspath(manifest_path))
181
-
182
- wer = wer_metric.compute(references=all_data['references'], predictions=predictions)
183
- wer = round(100 * wer, 2)
184
-
185
- # transcription_time = sum(all_results["transcription_time"])
186
- audio_length = sum(all_data["durations"])
187
- rtfx = audio_length / total_time
188
- rtfx = round(rtfx, 2)
189
-
190
- print("RTFX:", rtfx)
191
- print("WER:", wer, "%")
192
-
193
-
194
- if __name__ == "__main__":
195
- parser = argparse.ArgumentParser()
196
-
197
- parser.add_argument(
198
- "--model_id", type=str, required=True, help="Model identifier. Should be loadable with NVIDIA NeMo.",
199
- )
200
- parser.add_argument(
201
- '--dataset_path', type=str, default='hf-audio/open-asr-leaderboard', help='Dataset path. By default, it is `hf-audio/open-asr-leaderboard`'
202
- )
203
- parser.add_argument(
204
- '--data_cache_root', type=str, default=None, help='Root directory for audio cache. By default, it is the current working directory.'
205
- )
206
- parser.add_argument(
207
- "--dataset",
208
- type=str,
209
- required=True,
210
- help="Dataset name. *E.g.* `'librispeech_asr` for the LibriSpeech ASR dataset, or `'common_voice'` for Common Voice. The full list of dataset names "
211
- "can be found at `https://huggingface.co/datasets/hf-audio/open-asr-leaderboard`",
212
- )
213
- parser.add_argument(
214
- "--split",
215
- type=str,
216
- default="test",
217
- help="Split of the dataset. *E.g.* `'validation`' for the dev split, or `'test'` for the test split.",
218
- )
219
- parser.add_argument(
220
- "--device",
221
- type=int,
222
- default=-1,
223
- help="The device to run the pipeline on. -1 for CPU (default), 0 for the first GPU and so on.",
224
- )
225
- parser.add_argument(
226
- "--batch_size", type=int, default=32, help="Number of samples to go through each streamed batch.",
227
- )
228
- parser.add_argument(
229
- "--max_eval_samples",
230
- type=int,
231
- default=None,
232
- help="Number of samples to be evaluated. Put a lower number e.g. 64 for testing this script.",
233
- )
234
- parser.add_argument(
235
- "--streaming",
236
- action="store_true",
237
- help="Stream the dataset lazily over the network instead of downloading it in full before the evaluation. Off by default for reproducible benchmark timings.",
238
- )
239
- args = parser.parse_args()
240
-
241
- main(args)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
nemo_asr/run_eval_long.py DELETED
@@ -1,224 +0,0 @@
1
- import argparse
2
-
3
- import io
4
- import os
5
- import torch
6
- import evaluate
7
- import soundfile
8
-
9
- from tqdm import tqdm
10
- from normalizer import data_utils
11
- import numpy as np
12
-
13
- from nemo.collections.asr.models import ASRModel
14
- import time
15
-
16
-
17
- wer_metric = evaluate.load("wer")
18
-
19
-
20
- def main(args):
21
-
22
- DATA_CACHE_DIR = os.path.join(os.getcwd(), "audio_cache")
23
- DATASET_NAME = args.dataset
24
- SPLIT_NAME = args.split
25
-
26
- CACHE_DIR = os.path.join(DATA_CACHE_DIR, DATASET_NAME, SPLIT_NAME)
27
- if not os.path.exists(CACHE_DIR):
28
- os.makedirs(CACHE_DIR)
29
-
30
- if args.device >= 0:
31
- device = torch.device(f"cuda:{args.device}")
32
- compute_dtype=torch.bfloat16
33
- else:
34
- device = torch.device("cpu")
35
- compute_dtype=torch.float32
36
-
37
-
38
- if args.model_id.endswith(".nemo"):
39
- asr_model = ASRModel.restore_from(args.model_id, map_location=device)
40
- else:
41
- asr_model = ASRModel.from_pretrained(args.model_id, map_location=device) # type: ASRModel
42
-
43
- if args.longform:
44
- asr_model.change_attention_model("rel_pos_local_attn", [128, 128]) # local attn
45
- asr_model.to(compute_dtype)
46
- asr_model.eval()
47
- print(f"Model size: {sum(p.numel() for p in asr_model.parameters()) / 1e9:.2f}B parameters")
48
-
49
- dataset = data_utils.load_data(args)
50
-
51
- def download_audio_files(batch, indices):
52
-
53
- # download audio files and write the paths, transcriptions and durations to a manifest file
54
- audio_paths = []
55
- durations = []
56
-
57
- # Use global indices for unique filenames across all batches
58
- for global_idx, sample in zip(indices, batch["audio"]):
59
- # Use a unique filename based on global index
60
- audio_path = os.path.join(CACHE_DIR, f"sample_{global_idx}.wav")
61
-
62
- if "array" in sample:
63
- audio_array = np.float32(sample["array"])
64
- sample_rate = 16000
65
-
66
- elif "bytes" in sample: # added to be compatible with latest datasets library (3.x.x) that produces byte stream
67
- with io.BytesIO(sample["bytes"]) as audio_file:
68
- audio_array, sample_rate = soundfile.read(audio_file, dtype="float32")
69
-
70
- else:
71
- raise ValueError("Sample must have either 'array' or 'bytes' key")
72
-
73
- if not os.path.exists(audio_path):
74
- os.makedirs(os.path.dirname(audio_path), exist_ok=True)
75
- soundfile.write(audio_path, audio_array, sample_rate)
76
-
77
- audio_paths.append(audio_path)
78
- durations.append(len(audio_array) / sample_rate)
79
-
80
-
81
- batch["references"] = batch["norm_text"]
82
- batch["audio_filepaths"] = audio_paths
83
- batch["durations"] = durations
84
-
85
- return batch
86
-
87
-
88
- if args.max_eval_samples is not None and args.max_eval_samples > 0:
89
- print(f"Subsampling dataset to first {args.max_eval_samples} samples !")
90
- dataset = dataset.take(args.max_eval_samples)
91
-
92
- dataset = data_utils.prepare_data(dataset)
93
- if asr_model.cfg.decoding.strategy != "beam":
94
- asr_model.cfg.decoding.strategy = "greedy_batch"
95
- asr_model.change_decoding_strategy(asr_model.cfg.decoding)
96
-
97
- # prepraing the offline dataset
98
- dataset = dataset.map(download_audio_files, batch_size=args.batch_size, batched=True, with_indices=True, remove_columns=["audio"])
99
-
100
- # Write manifest from daraset batch using json and keys audio_filepath, duration, text
101
-
102
- all_data = {
103
- "audio_filepaths": [],
104
- "durations": [],
105
- "references": [],
106
- }
107
-
108
- data_itr = iter(dataset)
109
- for data in tqdm(data_itr, desc="Downloading Samples"):
110
- for key in all_data:
111
- all_data[key].append(data[key])
112
-
113
- # Sort audio_filepaths and references based on durations values
114
- sorted_indices = sorted(range(len(all_data["durations"])), key=lambda k: all_data["durations"][k], reverse=True)
115
- all_data["audio_filepaths"] = [all_data["audio_filepaths"][i] for i in sorted_indices]
116
- all_data["references"] = [all_data["references"][i] for i in sorted_indices]
117
- all_data["durations"] = [all_data["durations"][i] for i in sorted_indices]
118
-
119
- total_time = 0
120
- for _ in range(2): # warmup once and calculate rtf
121
- if _ == 0:
122
- audio_files = all_data["audio_filepaths"][:args.batch_size * 4] # warmup with 4 batches
123
- else:
124
- audio_files = all_data["audio_filepaths"]
125
- start_time = time.time()
126
- with torch.inference_mode(), torch.no_grad():
127
-
128
- if 'canary' in args.model_id and 'v2' not in args.model_id:
129
- pnc = 'nopnc'
130
- else:
131
- pnc = 'pnc'
132
-
133
- if 'canary' in args.model_id:
134
- transcriptions = asr_model.transcribe(audio_files, batch_size=args.batch_size, verbose=False, pnc=pnc, num_workers=1)
135
- else:
136
- transcriptions = asr_model.transcribe(audio_files, batch_size=args.batch_size, verbose=False, num_workers=1)
137
- end_time = time.time()
138
- if _ == 1:
139
- total_time += end_time - start_time
140
- total_time = total_time
141
-
142
- # normalize transcriptions with English normalizer
143
- if isinstance(transcriptions, tuple) and len(transcriptions) == 2:
144
- transcriptions = transcriptions[0]
145
- predictions = [data_utils.normalizer(pred.text) for pred in transcriptions]
146
-
147
- avg_time = total_time / len(all_data["audio_filepaths"])
148
-
149
- # Write manifest results (WER and RTFX)
150
- manifest_path = data_utils.write_manifest(
151
- all_data["references"],
152
- predictions,
153
- args.model_id,
154
- args.dataset_path,
155
- args.dataset,
156
- args.split,
157
- audio_length=all_data["durations"],
158
- transcription_time=[avg_time] * len(all_data["audio_filepaths"]),
159
- )
160
-
161
- print("Results saved at path:", os.path.abspath(manifest_path))
162
-
163
- wer = wer_metric.compute(references=all_data['references'], predictions=predictions)
164
- wer = round(100 * wer, 2)
165
-
166
- # transcription_time = sum(all_results["transcription_time"])
167
- audio_length = sum(all_data["durations"])
168
- rtfx = audio_length / total_time
169
- rtfx = round(rtfx, 2)
170
-
171
- print("RTFX:", rtfx)
172
- print("WER:", wer, "%")
173
-
174
-
175
- if __name__ == "__main__":
176
- parser = argparse.ArgumentParser()
177
-
178
- parser.add_argument(
179
- "--model_id", type=str, required=True, help="Model identifier. Should be loadable with NVIDIA NeMo.",
180
- )
181
- parser.add_argument(
182
- '--dataset_path', type=str, default='hf-audio/open-asr-leaderboard', help='Dataset path. By default, it is `hf-audio/open-asr-leaderboard`'
183
- )
184
- parser.add_argument(
185
- "--dataset",
186
- type=str,
187
- required=True,
188
- help="Dataset name. *E.g.* `'librispeech_asr` for the LibriSpeech ASR dataset, or `'common_voice'` for Common Voice. The full list of dataset names "
189
- "can be found at `https://huggingface.co/datasets/hf-audio/open-asr-leaderboard`",
190
- )
191
- parser.add_argument(
192
- "--split",
193
- type=str,
194
- default="test",
195
- help="Split of the dataset. *E.g.* `'validation`' for the dev split, or `'test'` for the test split.",
196
- )
197
- parser.add_argument(
198
- "--device",
199
- type=int,
200
- default=-1,
201
- help="The device to run the pipeline on. -1 for CPU (default), 0 for the first GPU and so on.",
202
- )
203
- parser.add_argument(
204
- "--batch_size", type=int, default=32, help="Number of samples to go through each streamed batch.",
205
- )
206
- parser.add_argument(
207
- "--max_eval_samples",
208
- type=int,
209
- default=None,
210
- help="Number of samples to be evaluated. Put a lower number e.g. 64 for testing this script.",
211
- )
212
- parser.add_argument(
213
- "--streaming",
214
- action="store_true",
215
- help="Stream the dataset lazily over the network instead of downloading it in full before the evaluation. Off by default for reproducible benchmark timings.",
216
- )
217
- parser.add_argument(
218
- "--longform",
219
- action="store_true",
220
- help="Whether to use longform mode.",
221
- )
222
- args = parser.parse_args()
223
-
224
- main(args)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
nemo_asr/run_eval_ml.py DELETED
@@ -1,280 +0,0 @@
1
- # This script is used to evaluate NeMo ASR models on the Multi-Lingual datasets
2
-
3
- import argparse
4
- import io
5
- import os
6
- os.environ["DATASETS_USE_TORCHCODEC"] = "0"
7
- import torch
8
- import evaluate
9
- import soundfile
10
- import numpy as np
11
- from tqdm import tqdm
12
- from datasets import load_dataset
13
- from normalizer import data_utils
14
- from normalizer.eval_utils import normalize_compound_pairs
15
- from nemo.collections.asr.models import ASRModel
16
- from omegaconf import OmegaConf
17
- import time
18
-
19
-
20
- wer_metric = evaluate.load("wer")
21
-
22
-
23
- def main(args):
24
- DATA_CACHE_DIR = os.path.join(os.getcwd(), "audio_cache")
25
- CONFIG_NAME = args.config_name
26
- SPLIT_NAME = args.split
27
-
28
- # Extract language from config_name if not provided
29
- if args.language:
30
- LANGUAGE = args.language
31
- else:
32
- # Extract language from config_name (e.g., "fleurs_en" -> "en")
33
- try:
34
- LANGUAGE = CONFIG_NAME.split('_', 1)[1]
35
- except IndexError:
36
- LANGUAGE = "en" # Default fallback
37
-
38
- print(f"Detected language: {LANGUAGE}")
39
-
40
- CACHE_DIR = os.path.join(DATA_CACHE_DIR, CONFIG_NAME, SPLIT_NAME)
41
- if not os.path.exists(CACHE_DIR):
42
- os.makedirs(CACHE_DIR)
43
-
44
- if args.device >= 0:
45
- device = torch.device(f"cuda:{args.device}")
46
- compute_dtype = torch.bfloat16
47
- else:
48
- device = torch.device("cpu")
49
- compute_dtype = torch.float32
50
-
51
- # Load ASR model
52
- if args.model_id.endswith(".nemo"):
53
- asr_model = ASRModel.restore_from(args.model_id, map_location=device)
54
- else:
55
- asr_model = ASRModel.from_pretrained(args.model_id, map_location=device)
56
-
57
- asr_model.to(compute_dtype)
58
- asr_model.eval()
59
- print(f"Model size: {sum(p.numel() for p in asr_model.parameters()) / 1e9:.2f}B parameters")
60
-
61
- # Load dataset using the HuggingFace dataset repository
62
- print(f"Loading dataset: {args.dataset} with config: {CONFIG_NAME}")
63
-
64
- dataset = load_dataset(args.dataset, CONFIG_NAME, split=SPLIT_NAME, streaming=args.streaming)
65
-
66
- if args.max_eval_samples is not None and args.max_eval_samples > 0:
67
- print(f"Subsampling dataset to first {args.max_eval_samples} samples!")
68
- dataset = dataset.select(range(min(args.max_eval_samples, len(dataset))))
69
-
70
- # Configure decoding strategy
71
- if asr_model.cfg.decoding.strategy != "beam":
72
- asr_model.cfg.decoding.strategy = "greedy_batch"
73
- if hasattr(asr_model.cfg.decoding, "greedy"):
74
- OmegaConf.update(asr_model.cfg.decoding, "greedy.use_cuda_graph_decoder", False, force_add=True)
75
- asr_model.change_decoding_strategy(asr_model.cfg.decoding)
76
-
77
- def download_audio_files(batch):
78
- """Process audio files and prepare them for evaluation."""
79
- audio_paths = []
80
- durations = []
81
-
82
- for i, (file_name, sample, duration, text) in enumerate(zip(
83
- batch["file_name"], batch["audio"], batch["duration"], batch["text"]
84
- )):
85
- # Create unique filename using index to avoid conflicts
86
- unique_id = f"{CONFIG_NAME}_{i}_{os.path.basename(file_name).replace('.wav', '')}"
87
- audio_path = os.path.join(CACHE_DIR, f"{unique_id}.wav")
88
-
89
- if "array" in sample:
90
- audio_array = np.float32(sample["array"])
91
- sample_rate = sample.get("sampling_rate", 16000)
92
- elif "bytes" in sample:
93
- with io.BytesIO(sample["bytes"]) as audio_file:
94
- audio_array, sample_rate = soundfile.read(audio_file, dtype="float32")
95
- else:
96
- raise ValueError("Sample must have either 'array' or 'bytes' key")
97
-
98
- if not os.path.exists(audio_path):
99
- os.makedirs(os.path.dirname(audio_path), exist_ok=True)
100
- soundfile.write(audio_path, audio_array, sample_rate)
101
-
102
- audio_paths.append(audio_path)
103
- # Use duration from dataset if available, otherwise calculate
104
- if duration is not None:
105
- durations.append(duration)
106
- else:
107
- durations.append(len(audio_array) / sample_rate)
108
-
109
- batch["references"] = [text for text in batch["text"]]
110
- batch["audio_filepaths"] = audio_paths
111
- batch["durations"] = durations
112
-
113
- return batch
114
-
115
- # Process the dataset
116
- print("Processing audio files...")
117
- dataset = dataset.map(
118
- download_audio_files,
119
- batch_size=args.batch_size,
120
- batched=True,
121
- remove_columns=["audio"]
122
- )
123
-
124
- # Collect all data
125
- all_data = {
126
- "audio_filepaths": [],
127
- "durations": [],
128
- "references": [],
129
- }
130
-
131
- print("Collecting data...")
132
- for data in tqdm(dataset, desc="Collecting samples"):
133
- all_data["audio_filepaths"].append(data["audio_filepaths"])
134
- all_data["durations"].append(data["durations"])
135
- all_data["references"].append(data["references"])
136
-
137
- # Sort by duration for efficient batch processing
138
- print("Sorting by duration...")
139
- sorted_indices = sorted(range(len(all_data["durations"])), key=lambda k: all_data["durations"][k], reverse=True)
140
- all_data["audio_filepaths"] = [all_data["audio_filepaths"][i] for i in sorted_indices]
141
- all_data["references"] = [all_data["references"][i] for i in sorted_indices]
142
- all_data["durations"] = [all_data["durations"][i] for i in sorted_indices]
143
-
144
- # Run evaluation with warmup
145
- total_time = 0
146
- for warmup_round in range(2): # warmup once and calculate rtf
147
- if warmup_round == 0:
148
- audio_files = all_data["audio_filepaths"][:args.batch_size * 4] # warmup with 4 batches
149
- print("Running warmup...")
150
- else:
151
- audio_files = all_data["audio_filepaths"]
152
- print("Running full evaluation...")
153
-
154
- start_time = time.time()
155
- with torch.inference_mode(), torch.no_grad():
156
- # for canary-1b and canary-1b-flash, we need to set pnc='no' for English and for other languages, we need to set pnc='pnc' but for canary-1b-v2 pnc='yes' for all languages
157
- if 'canary' in args.model_id and 'v2' not in args.model_id:
158
- pnc = 'nopnc' if LANGUAGE == "en" else 'pnc'
159
- else:
160
- pnc = 'pnc'
161
-
162
- if 'canary' in args.model_id:
163
- transcriptions = asr_model.transcribe(audio_files, batch_size=args.batch_size, verbose=False, pnc=pnc, num_workers=1, source_lang=LANGUAGE, target_lang=LANGUAGE)
164
- else:
165
- transcriptions = asr_model.transcribe(audio_files, batch_size=args.batch_size, verbose=False, num_workers=1)
166
- end_time = time.time()
167
-
168
- if warmup_round == 1:
169
- total_time = end_time - start_time
170
-
171
- # Process transcriptions
172
- if isinstance(transcriptions, tuple) and len(transcriptions) == 2:
173
- transcriptions = transcriptions[0]
174
-
175
- references = all_data["references"]
176
- if LANGUAGE == "en": # English is handled by the English normalizer
177
- references = [data_utils.normalizer(ref) for ref in references]
178
- predictions = [data_utils.normalizer(pred.text) for pred in transcriptions]
179
- else:
180
- references = [data_utils.ml_normalizer(ref, lang=LANGUAGE) for ref in references]
181
- predictions = [data_utils.ml_normalizer(pred.text, lang=LANGUAGE) for pred in transcriptions]
182
-
183
- # Filter empty references (consistent with English pipeline)
184
- filtered = [
185
- (ref, pred, dur)
186
- for ref, pred, dur in zip(references, predictions, all_data["durations"])
187
- if data_utils.is_target_text_in_range(ref)
188
- ]
189
- if filtered:
190
- references, predictions, all_data["durations"] = zip(*filtered)
191
- references, predictions = list(references), list(predictions)
192
- all_data["durations"] = list(all_data["durations"])
193
-
194
- avg_time = total_time / len(all_data["audio_filepaths"])
195
-
196
- # Write results using eval_utils.write_manifest
197
- manifest_path = data_utils.write_manifest(
198
- references,
199
- predictions,
200
- args.model_id,
201
- args.dataset, # dataset_path for filename
202
- CONFIG_NAME, # dataset_name
203
- SPLIT_NAME,
204
- audio_length=all_data["durations"],
205
- transcription_time=[avg_time] * len(all_data["audio_filepaths"]),
206
- )
207
-
208
- print("Results saved at path:", os.path.abspath(manifest_path))
209
-
210
- # Calculate metrics
211
- wer_refs, wer_preds = normalize_compound_pairs(references, predictions)
212
- wer = wer_metric.compute(references=wer_refs, predictions=wer_preds)
213
- wer = round(100 * wer, 2)
214
-
215
- audio_length = sum(all_data["durations"])
216
- rtfx = audio_length / total_time
217
- rtfx = round(rtfx, 2)
218
-
219
- print(f"Dataset: {args.dataset}")
220
- print(f"Language: {LANGUAGE}")
221
- print(f"Config: {CONFIG_NAME}")
222
- print(f"Model: {args.model_id}")
223
- print(f"RTFX: {rtfx}")
224
- print(f"WER: {wer}%")
225
-
226
-
227
- if __name__ == "__main__":
228
- parser = argparse.ArgumentParser()
229
-
230
- parser.add_argument(
231
- "--model_id", type=str, required=True, help="Model identifier. Should be loadable with NVIDIA NeMo.",
232
- )
233
- parser.add_argument(
234
- "--dataset",
235
- type=str,
236
- default="nithinraok/asr-leaderboard-datasets",
237
- help="Dataset name. Default is 'nithinraok/asr-leaderboard-datasets'"
238
- )
239
- parser.add_argument(
240
- "--config_name",
241
- type=str,
242
- required=True,
243
- help="Config name in format <dataset>_<lang> (e.g., fleurs_en, mcv_de, mls_es)"
244
- )
245
- parser.add_argument(
246
- "--language",
247
- type=str,
248
- default=None,
249
- help="Language code (e.g., en, de, es). If not provided, will be extracted from config_name."
250
- )
251
- parser.add_argument(
252
- "--split",
253
- type=str,
254
- default="test",
255
- help="Split of the dataset. Default is 'test'.",
256
- )
257
- parser.add_argument(
258
- "--device",
259
- type=int,
260
- default=-1,
261
- help="The device to run the pipeline on. -1 for CPU (default), 0 for the first GPU and so on.",
262
- )
263
- parser.add_argument(
264
- "--batch_size", type=int, default=32, help="Number of samples to go through each streamed batch.",
265
- )
266
- parser.add_argument(
267
- "--max_eval_samples",
268
- type=int,
269
- default=None,
270
- help="Number of samples to be evaluated. Put a lower number e.g. 64 for testing this script.",
271
- )
272
-
273
- parser.add_argument(
274
- "--streaming",
275
- action="store_true",
276
- help="Stream the dataset lazily over the network instead of downloading it in full before the evaluation. Off by default for reproducible benchmark timings.",
277
- )
278
- args = parser.parse_args()
279
-
280
- main(args)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
nemo_asr/run_eval_salm.py DELETED
@@ -1,276 +0,0 @@
1
- import argparse
2
-
3
- import io
4
- import os
5
- import torch
6
- import evaluate
7
- import soundfile
8
- import lhotse
9
-
10
- from tqdm import tqdm
11
- from normalizer import data_utils
12
- import numpy as np
13
-
14
- from nemo.collections.asr.models import ASRModel
15
- import time
16
-
17
-
18
- from nemo.collections.speechlm2.models.salm import SALM
19
- from omegaconf import OmegaConf
20
- from pathlib import Path
21
- from transformers import GenerationConfig
22
-
23
-
24
-
25
- wer_metric = evaluate.load("wer")
26
-
27
-
28
- class ToAudio(torch.utils.data.Dataset):
29
- def __getitem__(self, cuts):
30
- cuts = lhotse.CutSet([c.to_mono(mono_downmix=True) if isinstance(c, lhotse.MultiCut) else c for c in cuts])
31
- audios, audio_lens = cuts.load_audio(collate=True)
32
- return {"cuts": cuts, "audios": audios, "audio_lens": audio_lens}
33
-
34
-
35
- def setup_dloader(audio_files, batch_size, num_workers):
36
- cuts = lhotse.CutSet([lhotse.Recording.from_file(p).to_cut() for p in audio_files])
37
- cuts = cuts.resample(16000)
38
- return torch.utils.data.DataLoader(
39
- dataset=ToAudio(),
40
- sampler=lhotse.dataset.DynamicCutSampler(cuts, max_cuts=batch_size),
41
- num_workers=num_workers,
42
- batch_size=None,
43
- )
44
-
45
-
46
- def transcribe(model, dloader) -> list[str]:
47
- hyps = []
48
- eos_tokens = torch.tensor([model.text_eos_id])
49
- for batch_idx, batch in enumerate(dloader):
50
- answer_ids = model.generate(
51
- prompts=[
52
- [
53
- {"role": "user", "slots": {"message": f"Transcribe the following: {model.audio_locator_tag}"}}
54
- ]
55
- ] * len(batch["cuts"]),
56
- audios=batch["audios"].to(model.device, non_blocking=True),
57
- audio_lens=batch["audio_lens"].to(model.device, non_blocking=True),
58
- generation_config=GenerationConfig(
59
- max_new_tokens=128,
60
- bos_token_id=model.text_bos_id,
61
- eos_token_id=eos_tokens,
62
- pad_token_id=model.text_pad_id,
63
- ),
64
- )
65
- answer_ids = [parse_hyp(ans, eos_tokens) for ans in answer_ids.cpu()]
66
- hyps.extend(model.tokenizer.ids_to_text(ans).strip() for ans in answer_ids)
67
- return hyps
68
-
69
-
70
- def parse_hyp(answer: torch.Tensor, eos_tokens):
71
- end = (answer == torch.isin(answer, eos_tokens)).nonzero(as_tuple=True)[0]
72
- if end.numel() == 0:
73
- return answer
74
- end = end[0]
75
- return answer[:end]
76
-
77
-
78
- def main(args):
79
-
80
- data_cache_root = args.data_cache_root if args.data_cache_root is not None else os.getcwd()
81
- DATA_CACHE_DIR = os.path.join(data_cache_root, "audio_cache")
82
- DATASET_NAME = args.dataset
83
- SPLIT_NAME = args.split
84
-
85
- CACHE_DIR = os.path.join(DATA_CACHE_DIR, DATASET_NAME, SPLIT_NAME)
86
- if not os.path.exists(CACHE_DIR):
87
- os.makedirs(CACHE_DIR)
88
-
89
- torch.set_float32_matmul_precision("medium")
90
-
91
- device = torch.device(f"cuda:{args.device}")
92
- model = SALM.from_pretrained(args.model_id).eval().to(torch.bfloat16).to(device)
93
- print(f"Model size: {sum(p.numel() for p in model.parameters()) / 1e9:.2f}B parameters")
94
-
95
- dataset = data_utils.load_data(args)
96
-
97
- def download_audio_files(batch):
98
-
99
- # download audio files and write the paths, transcriptions and durations to a manifest file
100
- audio_paths = []
101
- original_audio_paths = []
102
- durations = []
103
- file_names = batch.get("file_name", [None] * len(batch["audio"]))
104
-
105
- # Use 'id' column if available, otherwise generate sequential IDs
106
- if "id" in batch:
107
- ids = batch["id"]
108
- else:
109
- # Generate IDs based on index
110
- start_idx = len([f for f in os.listdir(CACHE_DIR) if f.endswith('.wav')]) if os.path.exists(CACHE_DIR) else 0
111
- ids = [f"sample_{start_idx + i}" for i in range(len(batch["audio"]))]
112
-
113
- for id, file_name, sample in zip(ids, file_names, batch["audio"]):
114
-
115
- # first step added here to make ID and wav filenames unique
116
- # several datasets like earnings22 have a hierarchical structure
117
- # for eg. earnings22/test/4432298/281.wav, earnings22/test/4450488/281.wav
118
- # lhotse uses the filename (281.wav) here as unique ID to create and name cuts
119
- # ref: https://github.com/lhotse-speech/lhotse/blob/master/lhotse/dataset/collation.py#L186
120
- original_id = id # preserve before sanitization for use as audio_filepath
121
- id = id.replace('/', '_').removesuffix('.wav')
122
-
123
- audio_path = os.path.join(CACHE_DIR, f"{id}.wav")
124
- audio_array = np.float32(sample["array"])
125
- sample_rate = sample["sampling_rate"]
126
-
127
- if not os.path.exists(audio_path):
128
- os.makedirs(os.path.dirname(audio_path), exist_ok=True)
129
- soundfile.write(audio_path, audio_array, sample_rate)
130
-
131
- audio_paths.append(audio_path)
132
- if file_name is not None:
133
- original_audio_paths.append(os.path.basename(str(file_name)))
134
- else:
135
- original_audio_paths.append(original_id)
136
- durations.append(len(audio_array) / sample_rate)
137
-
138
-
139
- batch["references"] = batch["norm_text"]
140
- batch["audio_filepaths"] = audio_paths
141
- batch["original_audio_filepaths"] = original_audio_paths
142
- batch["durations"] = durations
143
-
144
- return batch
145
-
146
-
147
- if args.max_eval_samples is not None and args.max_eval_samples > 0:
148
- print(f"Subsampling dataset to first {args.max_eval_samples} samples !")
149
- dataset = dataset.take(args.max_eval_samples)
150
-
151
- dataset = data_utils.prepare_data(dataset)
152
-
153
- # prepraing the offline dataset
154
- dataset = dataset.map(download_audio_files, batch_size=args.batch_size, batched=True, remove_columns=["audio"])
155
-
156
- # Write manifest from daraset batch using json and keys audio_filepath, duration, text
157
-
158
- all_data = {
159
- "audio_filepaths": [],
160
- "original_audio_filepaths": [],
161
- "durations": [],
162
- "references": [],
163
- }
164
-
165
- data_itr = iter(dataset)
166
- for data in tqdm(data_itr, desc="Downloading Samples"):
167
- for key in all_data:
168
- all_data[key].append(data[key])
169
-
170
- # Sort audio_filepaths and references based on durations values
171
- sorted_indices = sorted(range(len(all_data["durations"])), key=lambda k: all_data["durations"][k], reverse=True)
172
- all_data["audio_filepaths"] = [all_data["audio_filepaths"][i] for i in sorted_indices]
173
- all_data["original_audio_filepaths"] = [all_data["original_audio_filepaths"][i] for i in sorted_indices]
174
- all_data["references"] = [all_data["references"][i] for i in sorted_indices]
175
- all_data["durations"] = [all_data["durations"][i] for i in sorted_indices]
176
-
177
-
178
- total_time = 0
179
- for _ in range(2): # warmup once and calculate rtf
180
- if _ == 0:
181
- audio_files = all_data["audio_filepaths"][:args.batch_size * 4] # warmup with 4 batches
182
- else:
183
- audio_files = all_data["audio_filepaths"]
184
- dloader = setup_dloader(audio_files=audio_files, batch_size=args.batch_size, num_workers=1)
185
- with torch.inference_mode():
186
- start_time = time.time()
187
- transcriptions = transcribe(model, dloader)
188
- end_time = time.time()
189
- if _ == 1:
190
- total_time += end_time - start_time
191
- total_time = total_time
192
-
193
- # normalize transcriptions with English normalizer
194
- if isinstance(transcriptions, tuple) and len(transcriptions) == 2:
195
- transcriptions = transcriptions[0]
196
- predictions = [data_utils.normalizer(pred) for pred in transcriptions]
197
-
198
- avg_time = total_time / len(all_data["audio_filepaths"])
199
-
200
- # Write manifest results (WER and RTFX)
201
- manifest_path = data_utils.write_manifest(
202
- all_data["references"],
203
- predictions,
204
- args.model_id,
205
- args.dataset_path,
206
- args.dataset,
207
- args.split,
208
- audio_length=all_data["durations"],
209
- transcription_time=[avg_time] * len(all_data["audio_filepaths"]),
210
- audio_filepaths=all_data["original_audio_filepaths"],
211
- )
212
-
213
- print("Results saved at path:", os.path.abspath(manifest_path))
214
-
215
- wer = wer_metric.compute(references=all_data['references'], predictions=predictions)
216
- wer = round(100 * wer, 2)
217
-
218
- audio_length = sum(all_data["durations"])
219
- rtfx = audio_length / total_time
220
- rtfx = round(rtfx, 2)
221
-
222
- print("RTFX:", rtfx)
223
- print("WER:", wer, "%")
224
-
225
-
226
- if __name__ == "__main__":
227
- parser = argparse.ArgumentParser()
228
-
229
- parser.add_argument(
230
- "--model_id", type=str, required=True, help="Model identifier. Should be loadable with NVIDIA NeMo.",
231
- )
232
- parser.add_argument(
233
- '--dataset_path', type=str, default='hf-audio/open-asr-leaderboard', help='Dataset path. By default, it is `hf-audio/open-asr-leaderboard`'
234
- )
235
- parser.add_argument(
236
- "--dataset",
237
- type=str,
238
- required=True,
239
- help="Dataset name. *E.g.* `'librispeech_asr` for the LibriSpeech ASR dataset, or `'common_voice'` for Common Voice. The full list of dataset names "
240
- "can be found at `https://huggingface.co/datasets/hf-audio/open-asr-leaderboard`",
241
- )
242
- parser.add_argument(
243
- "--split",
244
- type=str,
245
- default="test",
246
- help="Split of the dataset. *E.g.* `'validation`' for the dev split, or `'test'` for the test split.",
247
- )
248
- parser.add_argument(
249
- "--device",
250
- type=int,
251
- default=-1,
252
- help="The device to run the pipeline on. -1 for CPU (default), 0 for the first GPU and so on.",
253
- )
254
- parser.add_argument(
255
- "--batch_size", type=int, default=32, help="Number of samples to go through each streamed batch.",
256
- )
257
- parser.add_argument(
258
- "--max_eval_samples",
259
- type=int,
260
- default=None,
261
- help="Number of samples to be evaluated. Put a lower number e.g. 64 for testing this script.",
262
- )
263
- parser.add_argument(
264
- "--streaming",
265
- action="store_true",
266
- help="Stream the dataset lazily over the network instead of downloading it in full before the evaluation. Off by default for reproducible benchmark timings.",
267
- )
268
- parser.add_argument(
269
- "--data_cache_root",
270
- type=str,
271
- default=None,
272
- help="Root directory for caching audio files. Defaults to 'audio_cache' in current directory.",
273
- )
274
- args = parser.parse_args()
275
-
276
- main(args)