goavinash5/Gradio_LLAMA_Testing
0
1import os2import time3import argparse4 5from dotenv import load_dotenv6from distutils.util import strtobool7from memory_profiler import memory_usage8from tqdm import tqdm9 10from llama2_wrapper import LLAMA2_WRAPPER11 12 13def run_iteration(14 llama2_wrapper, prompt_example, DEFAULT_SYSTEM_PROMPT, DEFAULT_MAX_NEW_TOKENS15):16 def generation():17 generator = llama2_wrapper.run(18 prompt_example,19 [],20 DEFAULT_SYSTEM_PROMPT,21 DEFAULT_MAX_NEW_TOKENS,22 1,23 0.95,24 50,25 )26 model_response = None27 try:28 first_model_response = next(generator)29 except StopIteration:30 pass31 for model_response in generator:32 pass33 return llama2_wrapper.get_token_length(model_response), model_response34 35 tic = time.perf_counter()36 mem_usage, (output_token_length, model_response) = memory_usage(37 (generation,), max_usage=True, retval=True38 )39 toc = time.perf_counter()40 41 generation_time = toc - tic42 tokens_per_second = output_token_length / generation_time43 44 return generation_time, tokens_per_second, mem_usage, model_response45 46 47def main():48 parser = argparse.ArgumentParser()49 parser.add_argument("--iter", type=int, default=5, help="Number of iterations")50 parser.add_argument("--model_path", type=str, default="", help="model path")51 parser.add_argument(52 "--backend_type",53 type=str,54 default="",55 help="Backend options: llama.cpp, gptq, transformers",56 )57 parser.add_argument(58 "--load_in_8bit",59 type=bool,60 default=False,61 help="Whether to use bitsandbytes 8 bit.",62 )63 64 args = parser.parse_args()65 66 load_dotenv()67 68 DEFAULT_SYSTEM_PROMPT = os.getenv("DEFAULT_SYSTEM_PROMPT", "")69 MAX_MAX_NEW_TOKENS = int(os.getenv("MAX_MAX_NEW_TOKENS", 2048))70 DEFAULT_MAX_NEW_TOKENS = int(os.getenv("DEFAULT_MAX_NEW_TOKENS", 1024))71 MAX_INPUT_TOKEN_LENGTH = int(os.getenv("MAX_INPUT_TOKEN_LENGTH", 4000))72 73 MODEL_PATH = os.getenv("MODEL_PATH")74 assert MODEL_PATH is not None, f"MODEL_PATH is required, got: {MODEL_PATH}"75 BACKEND_TYPE = os.getenv("BACKEND_TYPE")76 assert BACKEND_TYPE is not None, f"BACKEND_TYPE is required, got: {BACKEND_TYPE}"77 78 LOAD_IN_8BIT = bool(strtobool(os.getenv("LOAD_IN_8BIT", "True")))79 80 if args.model_path != "":81 MODEL_PATH = args.model_path82 if args.backend_type != "":83 BACKEND_TYPE = args.backend_type84 if args.load_in_8bit:85 LOAD_IN_8BIT = True86 87 # Initialization88 init_tic = time.perf_counter()89 llama2_wrapper = LLAMA2_WRAPPER(90 model_path=MODEL_PATH,91 backend_type=BACKEND_TYPE,92 max_tokens=MAX_INPUT_TOKEN_LENGTH,93 load_in_8bit=LOAD_IN_8BIT,94 # verbose=True,95 )96 97 init_toc = time.perf_counter()98 initialization_time = init_toc - init_tic99 100 total_time = 0101 total_tokens_per_second = 0102 total_memory_gen = 0103 104 prompt_example = (105 "Can you explain briefly to me what is the Python programming language?"106 )107 108 # Cold run109 print("Performing cold run...")110 run_iteration(111 llama2_wrapper, prompt_example, DEFAULT_SYSTEM_PROMPT, DEFAULT_MAX_NEW_TOKENS112 )113 114 # Timed runs115 print(f"Performing {args.iter} timed runs...")116 for i in tqdm(range(args.iter)):117 try:118 gen_time, tokens_per_sec, mem_gen, model_response = run_iteration(119 llama2_wrapper,120 prompt_example,121 DEFAULT_SYSTEM_PROMPT,122 DEFAULT_MAX_NEW_TOKENS,123 )124 total_time += gen_time125 total_tokens_per_second += tokens_per_sec126 total_memory_gen += mem_gen127 except:128 break129 avg_time = total_time / (i + 1)130 avg_tokens_per_second = total_tokens_per_second / (i + 1)131 avg_memory_gen = total_memory_gen / (i + 1)132 133 print(f"Last model response: {model_response}")134 print(f"Initialization time: {initialization_time:0.4f} seconds.")135 print(136 f"Average generation time over {(i + 1)} iterations: {avg_time:0.4f} seconds."137 )138 print(139 f"Average speed over {(i + 1)} iterations: {avg_tokens_per_second:0.4f} tokens/sec."140 )141 print(f"Average memory usage during generation: {avg_memory_gen:.2f} MiB")142 143 144if __name__ == "__main__":145 main()146 