(self)
| 33 | super().__init__() |
| 34 | |
| 35 | def run(self) -> None: |
| 36 | main_model_logits_processor_list = LogitsProcessorList() |
| 37 | |
| 38 | processor_factory = ModelProcessorFactory() |
| 39 | |
| 40 | # 传递其他参数 |
| 41 | additional_kwargs = { |
| 42 | "learning_rate": self.learning_rate, |
| 43 | "ensemble_weight": self.information_dict['ensemble_weight'], |
| 44 | "learning_epochs_nums": self.learning_epochs_nums, |
| 45 | "ensemble_model_output_ids_queue": self.ensemble_model_output_ids_queue, |
| 46 | "assist_model_score_queue_list": self.assist_model_score_queue_list, |
| 47 | |
| 48 | "main_model_probability_transfer_matrix_list": self.main_model_probability_transfer_matrix_list, |
| 49 | "assist_model_probability_transfer_matrix_list": self.assist_model_probability_transfer_matrix_list, |
| 50 | "result_save_dir": self.result_save_dir, |
| 51 | "main_model_tokenizer": self.tokenizer, |
| 52 | "assist_model_tokenizer": self.assist_model_tokenizer, |
| 53 | "device": self.device, |
| 54 | "device_compute": self.device_compute, |
| 55 | "early_stop_string_list": self.early_stop_string_list, |
| 56 | } |
| 57 | # 创建对象 |
| 58 | logits_processor_mode = self.information_dict['logits_processor_mode'] |
| 59 | logits_processor_instance = processor_factory.create_processor(logits_processor_mode, |
| 60 | **additional_kwargs) |
| 61 | main_model_logits_processor_list.append(logits_processor_instance) |
| 62 | # main_model_logits_processor_list.append() |
| 63 | |
| 64 | main_model_input = self.information_dict['main_model_input'] |
| 65 | max_new_tokens = self.information_dict['max_new_tokens'] |
| 66 | main_model_input_ids = self.tokenizer(main_model_input, return_tensors="pt", |
| 67 | add_special_tokens=False).input_ids.to(self.device) |
| 68 | generation_kwargs = { |
| 69 | "input_ids": main_model_input_ids, |
| 70 | "max_new_tokens": max_new_tokens, |
| 71 | "do_sample": False, |
| 72 | "num_beams": 1, |
| 73 | "eos_token_id": self.tokenizer.eos_token_id, |
| 74 | "bos_token_id": self.tokenizer.bos_token_id, |
| 75 | # "pad_token_id": self.tokenizer.pad_token_id |
| 76 | } |
| 77 | |
| 78 | # generate_ids = self.model.generate(**generation_kwargs,pad_token_id=self.tokenizer.eos_token_id, |
| 79 | # logits_processor=main_model_logits_processor_list, |
| 80 | # streamer=self.model_streamer) |
| 81 | generate_ids = self.model.generate(**generation_kwargs, pad_token_id=self.tokenizer.eos_token_id, |
| 82 | logits_processor=main_model_logits_processor_list) |
| 83 | |
| 84 | text = self.tokenizer.decode(generate_ids[0]) |
| 85 | # print(text) |
| 86 | result_process_parameter = self.information_dict['result_process_parameter'] |
| 87 | split_key_before_list = result_process_parameter["split_key_before"] |
| 88 | split_key_behind_list = result_process_parameter["split_key_behind"] |
| 89 | |
| 90 | model_answer, prediction = answer_extract(text, self.information_dict['demon_count'], split_key_before_list, |
| 91 | split_key_behind_list) |
| 92 | print(self.information_dict['question']) |
nothing calls this directly
no test coverage detected