| import argparse |
| import sys |
| from time import perf_counter |
|
|
| import sys |
| |
| 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:" |
| |
|
|
| prompt = prompt_template.format(args.prompt) |
|
|
| test_batch_size = [1] |
| print("test_batch_size: ", test_batch_size) |
|
|
| for i, bs in enumerate(test_batch_size): |
| prompts = [prompt, ] * bs |
|
|
| |
| 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() |
|
|