lyraLLMs / lyrallms /LyraLlamaPy /examples /batch_stream_demo.py
carsonhxsu
# This is a combination of 22 commits.
8453337
Raw
History Blame Contribute Delete
4.6 kB
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 = "<human>:{}\n<bot>:" # 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()