| 16 | |
| 17 | |
| 18 | class Model: |
| 19 | def __init__(self, gpu=0): |
| 20 | self.inuse = False |
| 21 | print("Initializing...") |
| 22 | starting_time = time.time() |
| 23 | self.args = self.get_args() |
| 24 | self.pipeline = pipeline_runner(self.args, add_retrieval=False, server=True) |
| 25 | print("Loading model...") |
| 26 | self.llm = self.pipeline.get_backbone_model() |
| 27 | print("Model loaded in {} seconds".format(time.time() - starting_time)) |
| 28 | starting_time = time.time() |
| 29 | print("Loading retriever...") |
| 30 | self.retriever = self.pipeline.get_retriever() |
| 31 | print("Retriever loaded in {} seconds".format(time.time() - starting_time)) |
| 32 | self.query_id = 0 |
| 33 | # self.process_num = self.args.process_num |
| 34 | |
| 35 | self.queue = Queue() |
| 36 | self.callback = ServerEventCallback(self.queue) |
| 37 | self.occupied = False |
| 38 | |
| 39 | print("Server ready") |
| 40 | |
| 41 | def run_pipeline(self, user_input, method, top_k): |
| 42 | # empty the queue |
| 43 | while not self.queue.empty(): |
| 44 | self.queue.get() |
| 45 | self.query_id += 1 |
| 46 | temp_args = copy.deepcopy(self.args) |
| 47 | temp_args.retrieved_api_nums = top_k |
| 48 | temp_args.method = method |
| 49 | data_dict = { |
| 50 | "query": user_input, |
| 51 | } |
| 52 | self.pipeline.run_single_task( |
| 53 | method=method, |
| 54 | backbone_model=self.llm, |
| 55 | query_id=self.query_id, |
| 56 | data_dict=data_dict, |
| 57 | output_dir_path=self.args.output_answer_file, |
| 58 | retriever=self.retriever, |
| 59 | args=temp_args, |
| 60 | tool_des=None, |
| 61 | callbacks=[self.callback] |
| 62 | ) |
| 63 | |
| 64 | def get_queue(self): |
| 65 | while not self.queue.empty(): |
| 66 | yield self.queue.get() |
| 67 | |
| 68 | def get_args(self): |
| 69 | parser = argparse.ArgumentParser() |
| 70 | parser.add_argument('--corpus_tsv_path', type=str, default="your_retrival_corpus_path/", required=False, |
| 71 | help='') |
| 72 | parser.add_argument('--retrieval_model_path', type=str, default="your_model_path/", required=False, help='') |
| 73 | parser.add_argument('--retrieved_api_nums', type=int, default=5, required=False, help='') |
| 74 | parser.add_argument('--backbone_model', type=str, default="toolllama", required=False, |
| 75 | help='chatgpt_function or davinci or toolllama') |