| 65 | ] |
| 66 | |
| 67 | def __init__(self, |
| 68 | needle="", |
| 69 | haystack_file="", |
| 70 | retrieval_question="What are the special magic numbers for {}?", |
| 71 | results_version = 1, |
| 72 | rnd_number_digits = 7, |
| 73 | context_lengths_min = 1000, |
| 74 | context_lengths_max = 126000, |
| 75 | context_lengths_num_intervals = 10, |
| 76 | document_depth_percent_min = 0, |
| 77 | document_depth_percent_max = 100, |
| 78 | document_depth_percent_intervals = 10, |
| 79 | document_depth_percent_interval_type = "linear", |
| 80 | save_results = False, |
| 81 | final_context_length_buffer = 200, |
| 82 | print_ongoing_status = True): |
| 83 | needle="\nThe special magic {city} number is: {rnd_number}\n" |
| 84 | self.needle = needle |
| 85 | if not needle or not haystack_file or not retrieval_question: |
| 86 | raise ValueError("Needle, haystack, and retrieval_question must be provided.") |
| 87 | |
| 88 | self.rnd_number_digits = rnd_number_digits |
| 89 | self.context_lengths_num_intervals = context_lengths_num_intervals |
| 90 | self.document_depth_percent_intervals = document_depth_percent_intervals |
| 91 | self.haystack_file = haystack_file |
| 92 | self.retrieval_question = retrieval_question |
| 93 | self.results_version = results_version |
| 94 | self.save_results = save_results |
| 95 | self.final_context_length_buffer = final_context_length_buffer |
| 96 | self.print_ongoing_status = print_ongoing_status |
| 97 | self.testing_results = [] |
| 98 | |
| 99 | self.context_lengths = np.round(np.linspace(context_lengths_min, context_lengths_max, num=context_lengths_num_intervals, endpoint=True)).astype(int) |
| 100 | self.context_lengths = self.context_lengths.tolist() |
| 101 | if document_depth_percent_interval_type == 'linear': |
| 102 | self.document_depth_percents = np.round(np.linspace(document_depth_percent_min, document_depth_percent_max, num=document_depth_percent_intervals, endpoint=True)).astype(int) |
| 103 | elif document_depth_percent_interval_type == 'sigmoid': |
| 104 | self.document_depth_percents = [self.logistic(x) for x in np.linspace(document_depth_percent_min, document_depth_percent_max, document_depth_percent_intervals)] |
| 105 | else: |
| 106 | raise ValueError(f"Unsupported document_depth_percent_interval_type: {document_depth_percent_interval_type}") |
| 107 | self.document_depth_percents = self.document_depth_percents.tolist() |
| 108 | |
| 109 | self.model = Sampler() |
| 110 | |
| 111 | self.enc = AutoTokenizer.from_pretrained(FLAGS.tokenizer) |
| 112 | self.enc_tiktoken = tiktoken.encoding_for_model("gpt-4-1106-preview") |
| 113 | |
| 114 | def generate_random_number(self, num_digits): |
| 115 | lower_bound = 10**(num_digits - 1) |