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