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()