Optimized version using batch tokenization for better performance.
(self, tokenizer: PreTrainedTokenizerBase,
num_requests: int)
| 721 | random.shuffle(self.data) |
| 722 | |
| 723 | def sample(self, tokenizer: PreTrainedTokenizerBase, |
| 724 | num_requests: int) -> list[SampleRequest]: |
| 725 | """ |
| 726 | Optimized version using batch tokenization for better performance. |
| 727 | """ |
| 728 | # Collect all prompts and metadata |
| 729 | prompts = [] |
| 730 | max_tokens_list = [] |
| 731 | prompt_lengths = [] |
| 732 | |
| 733 | for i, entry in enumerate(self.data): |
| 734 | if len(prompts) >= num_requests: |
| 735 | break |
| 736 | prompt = entry["input"]["messages"][1]["content"] |
| 737 | max_tokens = entry["input"]["max_tokens"] |
| 738 | prompts.append(prompt) |
| 739 | max_tokens_list.append(max_tokens) |
| 740 | if "num_tokens" in entry["input"] and isinstance( |
| 741 | entry["input"]["num_tokens"], |
| 742 | int) and entry["input"]["num_tokens"] > 0: |
| 743 | prompt_lengths.append(entry["input"]["num_tokens"]) |
| 744 | |
| 745 | if len(prompt_lengths) > 0 and len(prompt_lengths) == len(prompts): |
| 746 | print( |
| 747 | f"skipping batch tokenization because prompt_lengths are already available" |
| 748 | ) |
| 749 | else: |
| 750 | prompt_lengths, _ = batch_tokenize_prompts( |
| 751 | prompts, tokenizer, progress_name="custom dataset prompts") |
| 752 | |
| 753 | # Create SampleRequest objects |
| 754 | samples = [] |
| 755 | for prompt, prompt_len, max_tokens in zip(prompts, prompt_lengths, |
| 756 | max_tokens_list): |
| 757 | samples.append( |
| 758 | SampleRequest( |
| 759 | prompt=prompt, |
| 760 | prompt_len=prompt_len, |
| 761 | expected_output_len=max_tokens, |
| 762 | )) |
| 763 | |
| 764 | return samples |
| 765 | |
| 766 | |
| 767 | # ----------------------------------------------------------------------------- |
nothing calls this directly
no test coverage detected