File size: 6,827 Bytes
88f4a67 3aea41a 88f4a67 3aea41a 88f4a67 3aea41a 88f4a67 ca7b658 88f4a67 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 | 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() |