MCPcopy Create free account
hub / github.com/PaddlePaddle/FastDeploy / test_sampler_logprobs

Function test_sampler_logprobs

tests/layers/test_sampler.py:201–223  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

199
200
201def test_sampler_logprobs():
202 batch_size = 32
203 vocab_size = 1024
204 min_seq_len = 1
205 max_seq_len = 1024
206 logprobs_mode_list = ["raw_logprobs", "raw_logits", "processed_logprobs", "processed_logits"]
207 logits = _create_fake_logits(batch_size, vocab_size)
208 sampling_metadata = _create_default_sampling_metadata(batch_size, min_seq_len, max_seq_len, max_num_logprobs=0)
209 for logprobs_mode in logprobs_mode_list:
210 fd_config = get_fd_config(batch_size)
211 fd_config.model_config.logprobs_mode = logprobs_mode
212 sampler = Sampler(logprobs_mode=logprobs_mode, fd_config=fd_config)
213 assert sampler.logprobs_mode == logprobs_mode
214 sampler_output = sampler(logits.clone(), sampling_metadata)
215 baseline_logprobs = get_baseline_logprobs(
216 logits.clone(), sampling_metadata, logprobs_mode=logprobs_mode, token_ids=sampler_output.sampled_token_ids
217 )
218 logprobs = sampler_output.logprobs_tensors.logprobs
219 print(f"baseline_logprobs = {baseline_logprobs}")
220 print(f"logprobs = {logprobs}")
221 equal = paddle.allclose(baseline_logprobs, logprobs, atol=1e-03, rtol=1e-03).item()
222 print(f"logprobs_mode: {logprobs_mode} equal={equal}")
223 assert equal
224
225
226if __name__ == "__main__":

Callers 1

test_sampler.pyFile · 0.85

Calls 7

SamplerClass · 0.90
get_baseline_logprobsFunction · 0.85
printFunction · 0.85
cloneMethod · 0.80
_create_fake_logitsFunction · 0.70
get_fd_configFunction · 0.70

Tested by

no test coverage detected