()
| 199 | |
| 200 | |
| 201 | def 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 | |
| 226 | if __name__ == "__main__": |
no test coverage detected