MCPcopy Create free account
hub / github.com/MLSysU/TD-Pipe / _sample

Function _sample

TD_Pipe/model_executor/layers/sampler.py:322–386  ·  view source on GitHub ↗
(
    probs: torch.Tensor,
    logprobs: torch.Tensor,
    sampling_metadata: SamplingMetadata,
)

Source from the content-addressed store, hash-verified

320
321
322def _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)

Callers 1

forwardMethod · 0.85

Calls 4

_multinomialFunction · 0.85
_greedy_sampleFunction · 0.85
_random_sampleFunction · 0.85
_beam_search_sampleFunction · 0.85

Tested by

no test coverage detected