newSpace / app.py
bailh
new
88f4a67
Raw
History Blame Contribute Delete
6.83 kB
import argparse
import random
import time
import numpy as np
import torch
import warnings
import gradio as gr
from transformers import AutoTokenizer, AutoModelForCausalLM
from model.model import MiniMindLM
from model.LMConfig import LMConfig
from model.model_lora import *
warnings.filterwarnings('ignore')
# 全局变量
model = None
tokenizer = None
last_prompt = ""
def init_model(args):
global model, tokenizer
tokenizer = AutoTokenizer.from_pretrained('./model/minimind_tokenizer')
print(f"BOS token: {tokenizer.bos_token}")
print(f"EOS token: {tokenizer.eos_token}")
print(f"PAD token: {tokenizer.pad_token}")
moe_path = '_moe' if args.use_moe else ''
modes = {0: 'pretrain', 1: 'full_sft', 2: 'rlhf', 3: 'reason'}
ckp = f'./{args.out_dir}/{modes[args.model_mode]}_{args.dim}{moe_path}.pth'
model = MiniMindLM(LMConfig(
dim=args.dim,
n_layers=args.n_layers,
max_seq_len=args.max_seq_len,
use_moe=args.use_moe
))
state_dict = torch.load(ckp, map_location=args.device)
model.load_state_dict({k: v for k, v in state_dict.items() if 'mask' not in k}, strict=True)
if args.lora_name != 'None':
apply_lora(model)
load_lora(model, f'./{args.out_dir}/lora/{args.lora_name}_{args.dim}.pth')
print(f'MiniMind模型参数量: {sum(p.numel() for p in model.parameters() if p.requires_grad) / 1e6:.2f}M(illion)')
return model.eval().to(args.device), tokenizer
# 设置可复现的随机种子
def setup_seed(seed):
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
def chat(prompt, history, args):
global last_prompt
last_prompt = prompt
try:
setup_seed(random.randint(0, 2048))
messages = history[-args.history_cnt:] if args.history_cnt else []
messages.append({"role": "user", "content": prompt})
new_prompt = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True
)[-args.max_seq_len + 1:] if args.model_mode != 0 else (tokenizer.bos_token + prompt)
print(f"Original prompt: {prompt}")
print(f"New prompt after processing: {new_prompt}")
# 检查 new_prompt 是否为空
if not new_prompt:
print("Warning: New prompt is empty.")
return history
input_ids = tokenizer(new_prompt)['input_ids']
print(f"Input IDs: {input_ids}")
if not input_ids:
print("Warning: Input IDs are empty.")
# 输出 tokenizer 配置信息,方便调试
print(f"Tokenizer config: {tokenizer}")
return history
answer = new_prompt
with torch.no_grad():
x = torch.tensor(input_ids, device=args.device).unsqueeze(0)
print(f"Input tensor x: {x}")
if x is None or x.numel() == 0:
print("Warning: Input tensor is empty.")
return history
outputs = model.generate(
x,
eos_token_id=tokenizer.eos_token_id,
max_new_tokens=args.max_seq_len,
temperature=args.temperature,
top_p=args.top_p,
stream=True,
pad_token_id=tokenizer.pad_token_id
)
history_idx = 0
for y in outputs:
# 完整显示标签
answer = tokenizer.decode(y[0].tolist(), skip_special_tokens=False)
if (answer and answer[-1] == '�') or not answer:
continue
answer = answer[history_idx:]
history_idx = len(answer)
yield history + [(prompt, answer)]
except Exception as e:
print(f"Error in chat function: {e}")
def repeat_last_prompt(key, msg_value):
global last_prompt
if key == "ArrowUp" and last_prompt:
return last_prompt
return msg_value
def default_question(history, args):
default_prompt = "介绍下自己。"
return default_prompt, history + [(default_prompt, None)]
def main():
parser = argparse.ArgumentParser(description="Chat with MiniMind")
parser.add_argument('--lora_name', default='None', type=str)
parser.add_argument('--out_dir', default='out', type=str)
parser.add_argument('--temperature', default=0.85, type=float)
parser.add_argument('--top_p', default=0.85, type=float)
parser.add_argument('--device', default='cuda' if torch.cuda.is_available() else 'cpu', type=str)
parser.add_argument('--dim', default=512, type=int)
parser.add_argument('--n_layers', default=8, type=int)
parser.add_argument('--max_seq_len', default=8192, type=int)
parser.add_argument('--use_moe', default=False, type=bool)
# 携带历史对话上下文条数
# history_cnt需要设为偶数,即【用户问题, 模型回答】为1组;设置为0时,即当前query不携带历史上文
# 模型未经过外推微调时,在更长的上下文的chat_template时难免出现性能的明显退化,因此需要注意此处设置
parser.add_argument('--history_cnt', default=0, type=int)
parser.add_argument('--stream', default=True, type=bool)
parser.add_argument('--load', default=0, type=int, help="0: 原生torch权重,1: transformers加载")
parser.add_argument('--model_mode', default=3, type=int,
help="0: 预训练模型,1: SFT-Chat模型,2: RLHF-Chat模型,3: Reason模型")
args = parser.parse_args()
global model, tokenizer
model, tokenizer = init_model(args)
with gr.Blocks() as demo:
chatbot = gr.Chatbot()
msg = gr.Textbox(value="介绍下自己") # 设置默认提问内容
ask_button = gr.Button("提问")
clear = gr.Button("Clear")
def user(user_message, history):
try:
if user_message.strip() == "":
return "", history
return user_message, history + [(user_message, None)]
except Exception as e:
print(f"Error in user function: {e}")
if hasattr(msg, 'keydown'):
msg.keydown(repeat_last_prompt, [gr.State(), msg], msg)
msg.submit(user, [msg, chatbot], [msg, chatbot], queue=False).then(
chat, [msg, chatbot, gr.State(args)], chatbot
)
ask_button.click(default_question, [chatbot, gr.State(args)], [msg, chatbot]).then(
chat, [msg, chatbot, gr.State(args)], chatbot
)
clear.click(lambda: None, None, chatbot, queue=False)
demo.queue()
demo.launch()
if __name__ == "__main__":
main()