MCPcopy Create free account
hub / github.com/NVIDIA/TensorRT-LLM / sample

Method sample

tensorrt_llm/serve/scripts/benchmark_dataset.py:1374–1410  ·  view source on GitHub ↗
(
        self,
        tokenizer: PreTrainedTokenizerBase,
        num_requests: int,
        output_len: Optional[int] = None,
        **kwargs,
    )

Source from the content-addressed store, hash-verified

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

Callers 9

sample_test_casesFunction · 0.45
sample_stagesFunction · 0.45
generate_samplesMethod · 0.45
_sample_loaded_dataMethod · 0.45
mainFunction · 0.45
_select_thoughtsMethod · 0.45

Calls 5

SampleRequestClass · 0.85
get_durationMethod · 0.80
appendMethod · 0.45
warningMethod · 0.45

Tested by 2

sample_test_casesFunction · 0.36
sample_stagesFunction · 0.36