(
probs: torch.Tensor,
logprobs: torch.Tensor,
sampling_metadata: SamplingMetadata,
)
| 320 | |
| 321 | |
| 322 | def _sample( |
| 323 | probs: torch.Tensor, |
| 324 | logprobs: torch.Tensor, |
| 325 | sampling_metadata: SamplingMetadata, |
| 326 | ) -> List[Tuple[List[int], List[int]]]: |
| 327 | categorized_seq_group_ids = {t: [] for t in SamplingType} |
| 328 | categorized_sample_indices = sampling_metadata.categorized_sample_indices |
| 329 | for i, seq_group in enumerate(sampling_metadata.seq_groups): |
| 330 | _, sampling_params = seq_group |
| 331 | sampling_type = sampling_params.sampling_type |
| 332 | categorized_seq_group_ids[sampling_type].append(i) |
| 333 | |
| 334 | sample_results_dict: Dict[int, Tuple[List[int], List[int]]] = {} |
| 335 | sample_metadata = {} |
| 336 | |
| 337 | # Counterintiutively, having two loops here is actually faster. |
| 338 | # The first loop can run without waiting on GPU<->CPU sync. |
| 339 | for sampling_type in SamplingType: |
| 340 | sample_indices = categorized_sample_indices[sampling_type] |
| 341 | num_tokens = len(sample_indices) |
| 342 | if num_tokens == 0: |
| 343 | continue |
| 344 | seq_group_ids = categorized_seq_group_ids[sampling_type] |
| 345 | seq_groups = [sampling_metadata.seq_groups[i] for i in seq_group_ids] |
| 346 | is_prompts = [i < sampling_metadata.num_prompts for i in seq_group_ids] |
| 347 | sample_metadata[sampling_type] = (seq_group_ids, seq_groups, |
| 348 | is_prompts, sample_indices) |
| 349 | if sampling_type == SamplingType.GREEDY: |
| 350 | greedy_samples = torch.argmax(logprobs[sample_indices], dim=-1) |
| 351 | elif sampling_type == SamplingType.RANDOM: |
| 352 | max_best_of = 1 |
| 353 | for seq_group, is_prompt in zip(seq_groups, is_prompts): |
| 354 | if is_prompt: |
| 355 | _, sampling_params = seq_group |
| 356 | max_best_of = max(max_best_of, sampling_params.best_of) |
| 357 | multinomial_samples = _multinomial(probs[sample_indices], |
| 358 | max_best_of) |
| 359 | elif sampling_type == SamplingType.BEAM: |
| 360 | beam_search_logprobs = logprobs[sample_indices] |
| 361 | else: |
| 362 | raise ValueError(f"Unsupported sampling type: {sampling_type}") |
| 363 | |
| 364 | # GPU<->CPU sync happens in the loop below. |
| 365 | |
| 366 | for sampling_type in SamplingType: |
| 367 | if sampling_type not in sample_metadata: |
| 368 | continue |
| 369 | seq_group_ids, seq_groups, is_prompts, sample_indices = sample_metadata[ |
| 370 | sampling_type] |
| 371 | if sampling_type == SamplingType.GREEDY: |
| 372 | sample_results = _greedy_sample(seq_groups, greedy_samples) |
| 373 | elif sampling_type == SamplingType.RANDOM: |
| 374 | sample_results = _random_sample(seq_groups, is_prompts, |
| 375 | multinomial_samples) |
| 376 | elif sampling_type == SamplingType.BEAM: |
| 377 | sample_results = _beam_search_sample(seq_groups, is_prompts, |
| 378 | sampling_metadata.seq_data, |
| 379 | beam_search_logprobs) |
no test coverage detected