| 7 | |
| 8 | |
| 9 | class 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 |
nothing calls this directly
no outgoing calls
no test coverage detected