MCPcopy Create free account
hub / github.com/LargeWorldModel/LWM / __init__

Method __init__

scripts/eval_needle.py:64–107  ·  view source on GitHub ↗
(self,
                 needle="",
                 haystack_file="",
                 retrieval_question="What is the special magic {} number?",
                 results_version = 1,
                 rnd_number_digits = 7,
                 context_lengths_min = 1000,
                 context_lengths_max = 126000,
                 context_lengths_num_intervals = 10,
                 document_depth_percent_min = 0,
                 document_depth_percent_max = 100,
                 document_depth_percent_intervals = 10,
                 document_depth_percent_interval_type = "linear",
                 save_results = False,
                 final_context_length_buffer = 200,
                 print_ongoing_status = True)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 2

logisticMethod · 0.95
SamplerClass · 0.70

Tested by

no test coverage detected