| 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}") |
|
|
| |
| 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.") |
| |
| 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) |
| |
| |
| |
| 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() |