MCPcopy Create free account
hub / github.com/LukeBailey181/sgs / SamplingAlgorithmBase

Class SamplingAlgorithmBase

sgs/verification/prover/algorithms/base.py:9–62  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

7
8
9class SamplingAlgorithmBase(object):
10 def __init__(self, scheduler, tokenizer_path, process_print, cfg, **kwargs):
11 os.environ["TOKENIZERS_PARALLELISM"] = "false"
12 self.scheduler = scheduler
13 self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
14 self.process_print = process_print
15 self.cfg = cfg
16
17 self.max_tokens = cfg.max_tokens
18 self.few_shot_dataset = cfg.get("few_shot_dataset", None)
19 if self.few_shot_dataset is not None:
20 self.few_shot_dataset = load_jsonl_objects(self.few_shot_dataset)
21 self.few_shot_num = cfg.get("few_shot_num", 3)
22 self.few_shot_func = MODEL_FORMAT[cfg.mode]["few_shot"]
23 self.log_interval = cfg.get("log_interval", 32)
24
25 @property
26 def algorithm_name(self):
27 return self.__class__.__name__
28
29 def _post_sample_info(self, **kwargs):
30 return dict(
31 algorithm=self.algorithm_name,
32 datetime=get_datetime(),
33 **kwargs,
34 )
35
36 def _encode_length(self, code):
37 return len(self.tokenizer.encode(code))
38
39 def _preprocess_data(self, input_data):
40 if self.few_shot_dataset is None or self.few_shot_num == 0:
41 return input_data
42 return {
43 **input_data,
44 "_extra_header": "".join(
45 [
46 self.few_shot_func(self.few_shot_dataset[idx])
47 for idx in np.random.choice(
48 [
49 _idx
50 for _idx, _data in enumerate(self.few_shot_dataset)
51 if _data["name"] != input_data["name"]
52 ],
53 size=self.few_shot_num,
54 replace=False,
55 )
56 ]
57 + [input_data.get("_extra_header", str())]
58 ),
59 }
60
61 def sample(self, **kwargs):
62 raise NotImplementedError

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected