| 1372 | skip_long_audios: bool = True |
| 1373 | |
| 1374 | def sample( |
| 1375 | self, |
| 1376 | tokenizer: PreTrainedTokenizerBase, |
| 1377 | num_requests: int, |
| 1378 | output_len: Optional[int] = None, |
| 1379 | **kwargs, |
| 1380 | ) -> list: |
| 1381 | import librosa |
| 1382 | output_len = (output_len |
| 1383 | if output_len is not None else self.DEFAULT_OUTPUT_LEN) |
| 1384 | prompt = ASRDataset.TRANSCRIPTION_PREAMBLE |
| 1385 | prompt_len = len(tokenizer(prompt).input_ids) |
| 1386 | sampled_requests = [] |
| 1387 | skipped = 0 |
| 1388 | for item in self.data: |
| 1389 | if len(sampled_requests) >= num_requests: |
| 1390 | break |
| 1391 | audio = item["audio"] |
| 1392 | y, sr = audio["array"], audio["sampling_rate"] |
| 1393 | duration_s = librosa.get_duration(y=y, sr=sr) |
| 1394 | # Whisper max supported duration |
| 1395 | if self.skip_long_audios and duration_s > 30: |
| 1396 | skipped += 1 |
| 1397 | continue |
| 1398 | |
| 1399 | sampled_requests.append( |
| 1400 | SampleRequest( |
| 1401 | prompt=prompt, |
| 1402 | prompt_len=prompt_len, |
| 1403 | expected_output_len=output_len, |
| 1404 | )) |
| 1405 | if skipped: |
| 1406 | logger.warning("%d samples discarded from dataset due to" \ |
| 1407 | " their length being greater than" \ |
| 1408 | " what Whisper supports.", skipped) |
| 1409 | self.maybe_oversample_requests(sampled_requests, num_requests) |
| 1410 | return sampled_requests |