import argparse import sys from time import perf_counter import sys # import ipdb sys.path.append('../') import threading import time from lyra_llama import lyraLlama def print_string(string, prev_seq_length=None, finish=False): if finish: print_list([string]) return print("\033c", end="") if prev_seq_length: print(string[:prev_seq_length], end='', flush=True) string = string[prev_seq_length:] for c_char in string: print(c_char, end='', flush=True) time.sleep(0.025) # 控制每个字符的输出间隔,可以根据需要调整 def print_list(lines): # 清空终端输出 print("\033c", end="") # 逐行打印字符串列表 print('\n'.join(lines)) def get_args(): parser = argparse.ArgumentParser(description="Faster ChatGLM6B Demo") parser.add_argument('--model-path', type=str, required=True, help='Model Path, include config.ini and tokenizer files') parser.add_argument('--tokenizer-path', type=str, default=None) parser.add_argument( '--data-type', type=str, metavar='TYPE', default='fp16', choices=[None, 'fp32', 'fp16', 'bf16', 'int8'], help='The data type to inference. If None, the data type follows the ' 'checkpoint data type.') parser.add_argument( '--memopt_mode', type=int, default=0, choices=[0, 1], help='Use MEMOPT mode to increase speed and reduce VRAM usage.' ' 0: FP16 mode' ' 1: Use MEMOPT mode') parser.add_argument( '--quant-type', type=str, metavar='TYPE', default='int8', choices=['int4', 'int8'], help='The data type of quantization. Only used in MEMOPT.') parser.add_argument( '--kvqparams-fpath', type=str, required=False, default="", help='File path of kv quantized params.') parser.add_argument("--prompt", type=str, required=False) parser.add_argument("--max-output-length", type=int, default=512) parser.add_argument("--warmups", type=int, default=10) parser.add_argument("--avgnums", type=int, default=10) args = parser.parse_args() print('\n=================== Arguments ===================') for k, v in vars(args).items(): print(f' - {k.ljust(25, ".")}: {v}') print('=================================================') return args def main(): args = get_args() model = lyraLlama(args.model_path, args.tokenizer_path, args.data_type, args.memopt_mode, args.quant_type, args.kvqparams_fpath) prompt_template = "Human: {}\n\nAssistant:" # xverse # prompt_template = ":{}\n:" # llama-ziya 13b prompt = prompt_template.format(args.prompt) test_batch_size = [1] # 8, 16, 32, 64 print("test_batch_size: ", test_batch_size) for i, bs in enumerate(test_batch_size): prompts = [prompt, ] * bs # warmup gpu for _ in range(args.warmups): for finish, output_texts in model.stream_generate(prompts, output_length=args.max_output_length, top_k=30, top_p=0.85, temperature=1.0, repetition_penalty=1.0, do_sample=False): pass start = perf_counter() for _ in range(args.avgnums): prev_sequence_lengths = None stream_counter = 0 for finish, output_texts in model.stream_generate(prompts, output_length=args.max_output_length, top_k=30, top_p=0.85, temperature=1.0, repetition_penalty=1.0, do_sample=False): if len(output_texts) == 1: print_string(output_texts[0], prev_sequence_lengths, finish) prev_sequence_lengths = len(output_texts[0]) else: print_list(output_texts) stream_counter += 1 end = perf_counter() cost = (end - start) / args.avgnums input_output_texts = [prompt + ' ' + gtext for prompt, gtext in zip(prompts, output_texts)] tokens = 0 input_tokens = len(model.tokenizer.encode(prompt)) words = 0 for text in input_output_texts: tokens += len(model.tokenizer.encode(text)) words += len(text) avg_output_tokens = tokens / len(input_output_texts) - input_tokens if __name__ == "__main__": main()