mirror of
https://github.com/jingyaogong/minimind.git
synced 2026-10-06 09:07:31 +00:00
250426
This commit is contained in:
@@ -1,8 +1,8 @@
|
||||
from openai import OpenAI
|
||||
|
||||
client = OpenAI(
|
||||
api_key="none",
|
||||
base_url="http://localhost:8998/v1"
|
||||
api_key="ollama",
|
||||
base_url="http://127.0.0.1:8998/v1"
|
||||
)
|
||||
stream = True
|
||||
conversation_history_origin = []
|
||||
|
||||
+46
-34
@@ -1,33 +1,57 @@
|
||||
import torch
|
||||
import warnings
|
||||
import sys
|
||||
import os
|
||||
import sys
|
||||
|
||||
__package__ = "scripts"
|
||||
sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), '..')))
|
||||
from transformers import AutoTokenizer, AutoModelForCausalLM
|
||||
from model.LMConfig import LMConfig
|
||||
from model.model import MiniMindLM
|
||||
import torch
|
||||
import warnings
|
||||
from transformers import AutoTokenizer, AutoModelForCausalLM, LlamaConfig, LlamaForCausalLM
|
||||
from model.model_minimind import MiniMindConfig, MiniMindForCausalLM
|
||||
|
||||
warnings.filterwarnings('ignore', category=UserWarning)
|
||||
|
||||
|
||||
def convert_torch2transformers(torch_path, transformers_path):
|
||||
def export_tokenizer(transformers_path):
|
||||
tokenizer = AutoTokenizer.from_pretrained('../model/minimind_tokenizer')
|
||||
tokenizer.save_pretrained(transformers_path)
|
||||
|
||||
LMConfig.register_for_auto_class()
|
||||
MiniMindLM.register_for_auto_class("AutoModelForCausalLM")
|
||||
lm_model = MiniMindLM(lm_config)
|
||||
# MoE模型需使用此函数转换
|
||||
def convert_torch2transformers_minimind(torch_path, transformers_path, dtype=torch.bfloat16):
|
||||
MiniMindConfig.register_for_auto_class()
|
||||
MiniMindForCausalLM.register_for_auto_class("AutoModelForCausalLM")
|
||||
lm_model = MiniMindForCausalLM(lm_config)
|
||||
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
||||
state_dict = torch.load(torch_path, map_location=device)
|
||||
lm_model.load_state_dict(state_dict, strict=False)
|
||||
lm_model = lm_model.to(dtype) # 转换模型权重精度
|
||||
model_params = sum(p.numel() for p in lm_model.parameters() if p.requires_grad)
|
||||
print(f'模型参数: {model_params / 1e6} 百万 = {model_params / 1e9} B (Billion)')
|
||||
lm_model.save_pretrained(transformers_path, safe_serialization=False)
|
||||
export_tokenizer(transformers_path)
|
||||
print(f"模型已保存为 Transformers 格式: {transformers_path}")
|
||||
tokenizer = AutoTokenizer.from_pretrained('../model/')
|
||||
tokenizer.save_pretrained(transformers_path)
|
||||
print(f"模型已保存为 Transformers-MiniMind 格式: {transformers_path}")
|
||||
|
||||
|
||||
# LlamaForCausalLM结构兼容第三方生态
|
||||
def convert_torch2transformers_llama(torch_path, transformers_path, dtype=torch.bfloat16):
|
||||
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
||||
state_dict = torch.load(torch_path, map_location=device)
|
||||
llama_config = LlamaConfig(
|
||||
vocab_size=lm_config.vocab_size,
|
||||
hidden_size=lm_config.hidden_size,
|
||||
intermediate_size=64 * ((int(lm_config.hidden_size * 8 / 3) + 64 - 1) // 64),
|
||||
num_hidden_layers=lm_config.num_hidden_layers,
|
||||
num_attention_heads=lm_config.num_attention_heads,
|
||||
num_key_value_heads=lm_config.num_key_value_heads,
|
||||
max_position_embeddings=lm_config.max_seq_len,
|
||||
rms_norm_eps=lm_config.rms_norm_eps,
|
||||
rope_theta=lm_config.rope_theta,
|
||||
)
|
||||
llama_model = LlamaForCausalLM(llama_config)
|
||||
llama_model.load_state_dict(state_dict, strict=False)
|
||||
llama_model = llama_model.to(dtype) # 转换模型权重精度
|
||||
llama_model.save_pretrained(transformers_path)
|
||||
model_params = sum(p.numel() for p in llama_model.parameters() if p.requires_grad)
|
||||
print(f'模型参数: {model_params / 1e6} 百万 = {model_params / 1e9} B (Billion)')
|
||||
tokenizer = AutoTokenizer.from_pretrained('../model/')
|
||||
tokenizer.save_pretrained(transformers_path)
|
||||
print(f"模型已保存为 Transformers-Llama 格式: {transformers_path}")
|
||||
|
||||
|
||||
def convert_transformers2torch(transformers_path, torch_path):
|
||||
@@ -36,27 +60,15 @@ def convert_transformers2torch(transformers_path, torch_path):
|
||||
print(f"模型已保存为 PyTorch 格式: {torch_path}")
|
||||
|
||||
|
||||
# don't need to use
|
||||
def push_to_hf(export_model_path):
|
||||
def init_model():
|
||||
tokenizer = AutoTokenizer.from_pretrained('../model/minimind_tokenizer')
|
||||
model = AutoModelForCausalLM.from_pretrained(export_model_path, trust_remote_code=True)
|
||||
return model, tokenizer
|
||||
|
||||
model, tokenizer = init_model()
|
||||
# model.push_to_hub(model_path)
|
||||
# tokenizer.push_to_hub(model_path, safe_serialization=False)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
lm_config = LMConfig(dim=512, n_layers=8, max_seq_len=8192, use_moe=False)
|
||||
lm_config = MiniMindConfig(hidden_size=768, num_hidden_layers=16, max_seq_len=8192, use_moe=True)
|
||||
|
||||
torch_path = f"../out/rlhf_{lm_config.dim}{'_moe' if lm_config.use_moe else ''}.pth"
|
||||
torch_path = f"../out/full_sft_{lm_config.hidden_size}{'_moe' if lm_config.use_moe else ''}.pth"
|
||||
|
||||
transformers_path = '../MiniMind2-Small'
|
||||
transformers_path = '../MiniMind2-MoE'
|
||||
|
||||
# convert torch to transformers model
|
||||
convert_torch2transformers(torch_path, transformers_path)
|
||||
convert_torch2transformers_minimind(torch_path, transformers_path)
|
||||
|
||||
# # convert transformers to torch model
|
||||
# convert_transformers2torch(transformers_path, torch_path)
|
||||
# # # convert transformers to torch model
|
||||
# # convert_transformers2torch(transformers_path, torch_path)
|
||||
|
||||
+74
-61
@@ -9,12 +9,14 @@ import time
|
||||
import torch
|
||||
import warnings
|
||||
import uvicorn
|
||||
|
||||
from threading import Thread
|
||||
from queue import Queue
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel
|
||||
from transformers import AutoTokenizer, AutoModelForCausalLM
|
||||
from model.LMConfig import LMConfig
|
||||
from model.model import MiniMindLM
|
||||
from transformers import AutoTokenizer, AutoModelForCausalLM, TextStreamer
|
||||
from model.model_minimind import MiniMindConfig, MiniMindForCausalLM
|
||||
from model.model_lora import apply_lora, load_lora
|
||||
|
||||
warnings.filterwarnings('ignore')
|
||||
@@ -23,30 +25,25 @@ app = FastAPI()
|
||||
|
||||
|
||||
def init_model(args):
|
||||
tokenizer = AutoTokenizer.from_pretrained('../model/minimind_tokenizer')
|
||||
if args.load == 0:
|
||||
tokenizer = AutoTokenizer.from_pretrained('../model/')
|
||||
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,
|
||||
ckp = f'../{args.out_dir}/{modes[args.model_mode]}_{args.hidden_size}{moe_path}.pth'
|
||||
model = MiniMindForCausalLM(MiniMindConfig(
|
||||
hidden_size=args.hidden_size,
|
||||
num_hidden_layers=args.num_hidden_layers,
|
||||
max_seq_len=args.max_seq_len,
|
||||
use_moe=args.use_moe
|
||||
))
|
||||
|
||||
state_dict = torch.load(ckp, map_location=device)
|
||||
model.load_state_dict({k: v for k, v in state_dict.items() if 'mask' not in k}, strict=True)
|
||||
|
||||
model.load_state_dict(torch.load(ckp, map_location=device), strict=True)
|
||||
if args.lora_name != 'None':
|
||||
apply_lora(model)
|
||||
load_lora(model, f'../{args.out_dir}/{args.lora_name}_{args.dim}.pth')
|
||||
load_lora(model, f'../{args.out_dir}/{args.lora_name}_{args.hidden_size}.pth')
|
||||
else:
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
'./MiniMind2',
|
||||
trust_remote_code=True
|
||||
)
|
||||
model_path = '../MiniMind2'
|
||||
model = AutoModelForCausalLM.from_pretrained(model_path, trust_remote_code=True)
|
||||
tokenizer = AutoTokenizer.from_pretrained(model_path)
|
||||
print(f'MiniMind模型参数量: {sum(p.numel() for p in model.parameters() if p.requires_grad) / 1e6:.2f}M(illion)')
|
||||
return model.eval().to(device), tokenizer
|
||||
|
||||
@@ -58,42 +55,61 @@ class ChatRequest(BaseModel):
|
||||
top_p: float = 0.92
|
||||
max_tokens: int = 8192
|
||||
stream: bool = False
|
||||
tools: list = []
|
||||
|
||||
|
||||
class CustomStreamer(TextStreamer):
|
||||
def __init__(self, tokenizer, queue):
|
||||
super().__init__(tokenizer, skip_prompt=True, skip_special_tokens=True)
|
||||
self.queue = queue
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
def on_finalized_text(self, text: str, stream_end: bool = False):
|
||||
self.queue.put(text)
|
||||
if stream_end:
|
||||
self.queue.put(None)
|
||||
|
||||
|
||||
def generate_stream_response(messages, temperature, top_p, max_tokens):
|
||||
try:
|
||||
new_prompt = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)[-max_tokens:]
|
||||
x = tokenizer(new_prompt).data['input_ids']
|
||||
x = (torch.tensor(x, dtype=torch.long, device=device)[None, ...])
|
||||
with torch.no_grad():
|
||||
res_y = model.generate(
|
||||
x,
|
||||
eos_token_id=tokenizer.eos_token_id,
|
||||
inputs = tokenizer(new_prompt, return_tensors="pt", truncation=True).to(device)
|
||||
|
||||
queue = Queue()
|
||||
streamer = CustomStreamer(tokenizer, queue)
|
||||
|
||||
def _generate():
|
||||
model.generate(
|
||||
inputs.input_ids,
|
||||
max_new_tokens=max_tokens,
|
||||
do_sample=True,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
stream=True,
|
||||
rp=1.,
|
||||
pad_token_id=tokenizer.pad_token_id
|
||||
attention_mask=inputs.attention_mask,
|
||||
pad_token_id=tokenizer.pad_token_id,
|
||||
eos_token_id=tokenizer.eos_token_id,
|
||||
streamer=streamer
|
||||
)
|
||||
history_idx = 0
|
||||
for y in res_y:
|
||||
answer = tokenizer.decode(y[0].tolist(), skip_special_tokens=True)
|
||||
if (answer and answer[-1] == '�') or not answer:
|
||||
continue
|
||||
delta = answer[history_idx:]
|
||||
history_idx = len(answer)
|
||||
json_data = {
|
||||
'id': f'chatcmpl-{int(time.time())}',
|
||||
'object': 'chat.completion.chunk',
|
||||
'created': int(time.time()),
|
||||
'model': 'minimind',
|
||||
'choices': [{'index': 0, 'delta': {'content': delta}, 'finish_reason': None}]
|
||||
}
|
||||
yield f"data: {json.dumps(json_data)}\n\n"
|
||||
|
||||
Thread(target=_generate).start()
|
||||
|
||||
while True:
|
||||
text = queue.get()
|
||||
if text is None:
|
||||
yield json.dumps({
|
||||
"choices": [{
|
||||
"delta": {},
|
||||
"finish_reason": "stop"
|
||||
}]
|
||||
}, ensure_ascii=False)
|
||||
break
|
||||
|
||||
yield json.dumps({
|
||||
"choices": [{"delta": {"content": text}}]
|
||||
}, ensure_ascii=False)
|
||||
|
||||
except Exception as e:
|
||||
yield f"data: {json.dumps({'error': str(e)})}\n\n"
|
||||
yield json.dumps({"error": str(e)})
|
||||
|
||||
|
||||
@app.post("/v1/chat/completions")
|
||||
@@ -101,12 +117,12 @@ async def chat_completions(request: ChatRequest):
|
||||
try:
|
||||
if request.stream:
|
||||
return StreamingResponse(
|
||||
generate_stream_response(
|
||||
(f"data: {chunk}\n\n" for chunk in generate_stream_response(
|
||||
messages=request.messages,
|
||||
temperature=request.temperature,
|
||||
top_p=request.top_p,
|
||||
max_tokens=request.max_tokens
|
||||
),
|
||||
)),
|
||||
media_type="text/event-stream"
|
||||
)
|
||||
else:
|
||||
@@ -115,20 +131,19 @@ async def chat_completions(request: ChatRequest):
|
||||
tokenize=False,
|
||||
add_generation_prompt=True
|
||||
)[-request.max_tokens:]
|
||||
x = tokenizer(new_prompt).data['input_ids']
|
||||
x = (torch.tensor(x, dtype=torch.long, device=device)[None, ...])
|
||||
inputs = tokenizer(new_prompt, return_tensors="pt", truncation=True).to(device)
|
||||
with torch.no_grad():
|
||||
res_y = model.generate(
|
||||
x,
|
||||
generated_ids = model.generate(
|
||||
inputs["input_ids"],
|
||||
max_length=inputs["input_ids"].shape[1] + request.max_tokens,
|
||||
do_sample=True,
|
||||
attention_mask=inputs["attention_mask"],
|
||||
pad_token_id=tokenizer.pad_token_id,
|
||||
eos_token_id=tokenizer.eos_token_id,
|
||||
max_new_tokens=request.max_tokens,
|
||||
temperature=request.temperature,
|
||||
top_p=request.top_p,
|
||||
stream=False,
|
||||
rp=1.,
|
||||
pad_token_id=tokenizer.pad_token_id
|
||||
temperature=request.temperature
|
||||
)
|
||||
answer = tokenizer.decode(res_y.squeeze()[x.shape[1]:].tolist(), skip_special_tokens=True)
|
||||
answer = tokenizer.decode(generated_ids[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True)
|
||||
return {
|
||||
"id": f"chatcmpl-{int(time.time())}",
|
||||
"object": "chat.completion",
|
||||
@@ -142,7 +157,6 @@ async def chat_completions(request: ChatRequest):
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@@ -151,14 +165,13 @@ if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="Server for MiniMind")
|
||||
parser.add_argument('--out_dir', default='out', type=str)
|
||||
parser.add_argument('--lora_name', default='None', type=str)
|
||||
parser.add_argument('--dim', default=512, type=int)
|
||||
parser.add_argument('--n_layers', default=8, type=int)
|
||||
parser.add_argument('--hidden_size', default=768, type=int)
|
||||
parser.add_argument('--num_hidden_layers', default=16, type=int)
|
||||
parser.add_argument('--max_seq_len', default=8192, type=int)
|
||||
parser.add_argument('--use_moe', default=False, type=bool)
|
||||
parser.add_argument('--load', default=0, type=int, help="0: 从原生torch权重,1: 利用transformers加载")
|
||||
parser.add_argument('--model_mode', default=1, type=int, help="0: 预训练模型,1: SFT-Chat模型,2: RLHF-Chat模型,3: Reason模型")
|
||||
|
||||
parser.add_argument('--model_mode', default=1, type=int,
|
||||
help="0: 预训练模型,1: SFT-Chat模型,2: RLHF-Chat模型,3: Reason模型")
|
||||
device = 'cuda' if torch.cuda.is_available() else 'cpu'
|
||||
model, tokenizer = init_model(parser.parse_args())
|
||||
|
||||
uvicorn.run(app, host="0.0.0.0", port=8998)
|
||||
|
||||
+15
-20
@@ -1,14 +1,9 @@
|
||||
import random
|
||||
from tqdm import tqdm
|
||||
from transformers import AutoTokenizer
|
||||
import json
|
||||
from datasets import load_dataset
|
||||
from tokenizers import (
|
||||
decoders,
|
||||
models,
|
||||
normalizers,
|
||||
pre_tokenizers,
|
||||
processors,
|
||||
trainers,
|
||||
Tokenizer,
|
||||
)
|
||||
@@ -32,7 +27,7 @@ def train_tokenizer():
|
||||
tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=False)
|
||||
|
||||
# 定义特殊token
|
||||
special_tokens = ["<unk>", "<s>", "</s>"]
|
||||
special_tokens = ["<|endoftext|>", "<|im_start|>", "<|im_end|>"]
|
||||
|
||||
# 设置训练器并添加特殊token
|
||||
trainer = trainers.BpeTrainer(
|
||||
@@ -52,15 +47,15 @@ def train_tokenizer():
|
||||
tokenizer.decoder = decoders.ByteLevel()
|
||||
|
||||
# 检查特殊token的索引
|
||||
assert tokenizer.token_to_id("<unk>") == 0
|
||||
assert tokenizer.token_to_id("<s>") == 1
|
||||
assert tokenizer.token_to_id("</s>") == 2
|
||||
assert tokenizer.token_to_id("<|endoftext|>") == 0
|
||||
assert tokenizer.token_to_id("<|im_start|>") == 1
|
||||
assert tokenizer.token_to_id("<|im_end|>") == 2
|
||||
|
||||
# 保存tokenizer
|
||||
tokenizer_dir = "../model/minimind_tokenizer"
|
||||
tokenizer_dir = "../model/"
|
||||
os.makedirs(tokenizer_dir, exist_ok=True)
|
||||
tokenizer.save(os.path.join(tokenizer_dir, "tokenizer.json"))
|
||||
tokenizer.model.save("../model/minimind_tokenizer")
|
||||
tokenizer.model.save("../model/")
|
||||
|
||||
# 手动创建配置文件
|
||||
config = {
|
||||
@@ -69,7 +64,7 @@ def train_tokenizer():
|
||||
"add_prefix_space": False,
|
||||
"added_tokens_decoder": {
|
||||
"0": {
|
||||
"content": "<unk>",
|
||||
"content": "<|endoftext|>",
|
||||
"lstrip": False,
|
||||
"normalized": False,
|
||||
"rstrip": False,
|
||||
@@ -77,7 +72,7 @@ def train_tokenizer():
|
||||
"special": True
|
||||
},
|
||||
"1": {
|
||||
"content": "<s>",
|
||||
"content": "<|im_start|>",
|
||||
"lstrip": False,
|
||||
"normalized": False,
|
||||
"rstrip": False,
|
||||
@@ -85,7 +80,7 @@ def train_tokenizer():
|
||||
"special": True
|
||||
},
|
||||
"2": {
|
||||
"content": "</s>",
|
||||
"content": "<|im_end|>",
|
||||
"lstrip": False,
|
||||
"normalized": False,
|
||||
"rstrip": False,
|
||||
@@ -94,17 +89,17 @@ def train_tokenizer():
|
||||
}
|
||||
},
|
||||
"additional_special_tokens": [],
|
||||
"bos_token": "<s>",
|
||||
"bos_token": "<|im_start|>",
|
||||
"clean_up_tokenization_spaces": False,
|
||||
"eos_token": "</s>",
|
||||
"eos_token": "<|im_end|>",
|
||||
"legacy": True,
|
||||
"model_max_length": 32768,
|
||||
"pad_token": "<unk>",
|
||||
"pad_token": "<|endoftext|>",
|
||||
"sp_model_kwargs": {},
|
||||
"spaces_between_special_tokens": False,
|
||||
"tokenizer_class": "PreTrainedTokenizerFast",
|
||||
"unk_token": "<unk>",
|
||||
"chat_template": "{% if messages[0]['role'] == 'system' %}{% set system_message = messages[0]['content'] %}{{ '<s>system\\n' + system_message + '</s>\\n' }}{% else %}{{ '<s>system\\n你是 MiniMind,是一个有用的人工智能助手。</s>\\n' }}{% endif %}{% for message in messages %}{% set content = message['content'] %}{% if message['role'] == 'user' %}{{ '<s>user\\n' + content + '</s>\\n<s>assistant\\n' }}{% elif message['role'] == 'assistant' %}{{ content + '</s>' + '\\n' }}{% endif %}{% endfor %}"
|
||||
"unk_token": "<|endoftext|>",
|
||||
"chat_template": "{% if messages[0]['role'] == 'system' %}{% set system_message = messages[0]['content'] %}{{ '<|im_start|>system\\n' + system_message + '<|im_end|>\\n' }}{% else %}{{ '<|im_start|>system\\nYou are a helpful assistant<|im_end|>\\n' }}{% endif %}{% for message in messages %}{% set content = message['content'] %}{% if message['role'] == 'user' %}{{ '<|im_start|>user\\n' + content + '<|im_end|>\\n<|im_start|>assistant\\n' }}{% elif message['role'] == 'assistant' %}{{ content + '<|im_end|>' + '\\n' }}{% endif %}{% endfor %}"
|
||||
}
|
||||
|
||||
# 保存配置文件
|
||||
@@ -118,7 +113,7 @@ def eval_tokenizer():
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
# 加载预训练的tokenizer
|
||||
tokenizer = AutoTokenizer.from_pretrained("../model/minimind_tokenizer")
|
||||
tokenizer = AutoTokenizer.from_pretrained("../model/")
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": "你是一个优秀的聊天机器人,总是给我正确的回应!"},
|
||||
|
||||
+100
-65
@@ -1,14 +1,13 @@
|
||||
import random
|
||||
import re
|
||||
import time
|
||||
from threading import Thread
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
import streamlit as st
|
||||
import torch
|
||||
|
||||
st.set_page_config(page_title="MiniMind", initial_sidebar_state="collapsed")
|
||||
|
||||
# 在文件开头的 CSS 样式中修改按钮样式
|
||||
st.markdown("""
|
||||
<style>
|
||||
/* 添加操作按钮样式 */
|
||||
@@ -70,7 +69,9 @@ device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
|
||||
def process_assistant_content(content):
|
||||
if 'R1' not in MODEL_PATHS[selected_model][1]:
|
||||
if model_source == "API" and 'R1' not in api_model_name:
|
||||
return content
|
||||
if model_source != "API" and 'R1' not in MODEL_PATHS[selected_model][1]:
|
||||
return content
|
||||
|
||||
if '<think>' in content and '</think>' in content:
|
||||
@@ -119,7 +120,6 @@ def init_chat_messages():
|
||||
if message["role"] == "assistant":
|
||||
with st.chat_message("assistant", avatar=image_url):
|
||||
st.markdown(process_assistant_content(message["content"]), unsafe_allow_html=True)
|
||||
# 在消息内容下方添加按钮
|
||||
if st.button("🗑", key=f"delete_{i}"):
|
||||
st.session_state.messages.pop(i)
|
||||
st.session_state.messages.pop(i - 1)
|
||||
@@ -137,8 +137,6 @@ def init_chat_messages():
|
||||
|
||||
return st.session_state.messages
|
||||
|
||||
|
||||
# 添加这两个辅助函数
|
||||
def regenerate_answer(index):
|
||||
st.session_state.messages.pop()
|
||||
st.session_state.chat_messages.pop()
|
||||
@@ -153,32 +151,34 @@ def delete_conversation(index):
|
||||
st.rerun()
|
||||
|
||||
|
||||
# 侧边栏模型选择
|
||||
st.sidebar.title("模型设定调整")
|
||||
|
||||
st.sidebar.text("【注】训练数据偏差,增加上下文记忆时\n多轮对话(较单轮)容易出现能力衰减")
|
||||
# st.sidebar.text("训练数据偏差,增加上下文记忆时\n多轮对话(较单轮)容易出现能力衰减")
|
||||
st.session_state.history_chat_num = st.sidebar.slider("Number of Historical Dialogues", 0, 6, 0, step=2)
|
||||
# st.session_state.history_chat_num = 0
|
||||
st.session_state.max_new_tokens = st.sidebar.slider("Max Sequence Length", 256, 8192, 8192, step=1)
|
||||
st.session_state.top_p = st.sidebar.slider("Top-P", 0.8, 0.99, 0.85, step=0.01)
|
||||
st.session_state.temperature = st.sidebar.slider("Temperature", 0.6, 1.2, 0.85, step=0.01)
|
||||
|
||||
# 模型路径映射
|
||||
MODEL_PATHS = {
|
||||
"MiniMind2-R1 (0.1B)": ["../MiniMind2-R1", "MiniMind2-R1"],
|
||||
"MiniMind2-Small-R1 (0.02B)": ["../MiniMind2-Small-R1", "MiniMind2-Small-R1"],
|
||||
"MiniMind2 (0.1B)": ["../MiniMind2", "MiniMind2"],
|
||||
"MiniMind2-MoE (0.15B)": ["../MiniMind2-MoE", "MiniMind2-MoE"],
|
||||
"MiniMind2-Small (0.02B)": ["../MiniMind2-Small", "MiniMind2-Small"],
|
||||
"MiniMind-V1 (0.1B)": ["../minimind-v1", "MiniMind-V1"],
|
||||
"MiniMind-V1-MoE (0.1B)": ["../minimind-v1-moe", "MiniMind-V1-MoE"],
|
||||
"MiniMind-V1-Small (0.02B)": ["../minimind-v1-small", "MiniMind-V1-Small"],
|
||||
}
|
||||
model_source = st.sidebar.radio("选择模型来源", ["本地模型", "API"], index=0)
|
||||
|
||||
selected_model = st.sidebar.selectbox('Models', list(MODEL_PATHS.keys()), index=2) # 默认选择 MiniMind2
|
||||
model_path = MODEL_PATHS[selected_model][0]
|
||||
if model_source == "API":
|
||||
api_url = st.sidebar.text_input("API URL", value="http://127.0.0.1:8000/v1")
|
||||
api_model_id = st.sidebar.text_input("Model ID", value="minimind")
|
||||
api_model_name = st.sidebar.text_input("Model Name", value="MiniMind2")
|
||||
api_key = st.sidebar.text_input("API Key", value="none", type="password")
|
||||
slogan = f"Hi, I'm {api_model_name}"
|
||||
else:
|
||||
MODEL_PATHS = {
|
||||
"MiniMind2-R1 (0.1B)": ["../MiniMind2-R1", "MiniMind2-R1"],
|
||||
"MiniMind2-Small-R1 (0.02B)": ["../MiniMind2-Small-R1", "MiniMind2-Small-R1"],
|
||||
"MiniMind2 (0.1B)": ["../MiniMind2", "MiniMind2"],
|
||||
"MiniMind2-MoE (0.15B)": ["../MiniMind2-MoE", "MiniMind2-MoE"],
|
||||
"MiniMind2-Small (0.02B)": ["../MiniMind2-Small", "MiniMind2-Small"]
|
||||
}
|
||||
|
||||
slogan = f"Hi, I'm {MODEL_PATHS[selected_model][1]}"
|
||||
selected_model = st.sidebar.selectbox('Models', list(MODEL_PATHS.keys()), index=2) # 默认选择 MiniMind2
|
||||
model_path = MODEL_PATHS[selected_model][0]
|
||||
slogan = f"Hi, I'm {MODEL_PATHS[selected_model][1]}"
|
||||
|
||||
image_url = "https://www.modelscope.cn/api/v1/studio/gongjy/MiniMind/repo?Revision=master&FilePath=images%2Flogo2.png&View=true"
|
||||
|
||||
@@ -205,23 +205,22 @@ def setup_seed(seed):
|
||||
|
||||
|
||||
def main():
|
||||
model, tokenizer = load_model_tokenizer(model_path)
|
||||
if model_source == "本地模型":
|
||||
model, tokenizer = load_model_tokenizer(model_path)
|
||||
else:
|
||||
model, tokenizer = None, None
|
||||
|
||||
# 初始化消息列表
|
||||
if "messages" not in st.session_state:
|
||||
st.session_state.messages = []
|
||||
st.session_state.chat_messages = []
|
||||
|
||||
# Use session state messages
|
||||
messages = st.session_state.messages
|
||||
|
||||
# 在显示历史消息的循环中
|
||||
for i, message in enumerate(messages):
|
||||
if message["role"] == "assistant":
|
||||
with st.chat_message("assistant", avatar=image_url):
|
||||
st.markdown(process_assistant_content(message["content"]), unsafe_allow_html=True)
|
||||
if st.button("×", key=f"delete_{i}"):
|
||||
# 删除当前消息及其之后的所有消息
|
||||
st.session_state.messages = st.session_state.messages[:i - 1]
|
||||
st.session_state.chat_messages = st.session_state.chat_messages[:i - 1]
|
||||
st.rerun()
|
||||
@@ -230,14 +229,11 @@ def main():
|
||||
f'<div style="display: flex; justify-content: flex-end;"><div style="display: inline-block; margin: 10px 0; padding: 8px 12px 8px 12px; background-color: gray; border-radius: 10px; color:white; ">{message["content"]}</div></div>',
|
||||
unsafe_allow_html=True)
|
||||
|
||||
# 处理新的输入或重新生成
|
||||
prompt = st.chat_input(key="input", placeholder="给 MiniMind 发送消息")
|
||||
|
||||
# 检查是否需要重新生成
|
||||
if hasattr(st.session_state, 'regenerate') and st.session_state.regenerate:
|
||||
prompt = st.session_state.last_user_message
|
||||
regenerate_index = st.session_state.regenerate_index # 获取重新生成的位置
|
||||
# 清除所有重新生成相关的状态
|
||||
regenerate_index = st.session_state.regenerate_index
|
||||
delattr(st.session_state, 'regenerate')
|
||||
delattr(st.session_state, 'last_user_message')
|
||||
delattr(st.session_state, 'regenerate_index')
|
||||
@@ -246,48 +242,87 @@ def main():
|
||||
st.markdown(
|
||||
f'<div style="display: flex; justify-content: flex-end;"><div style="display: inline-block; margin: 10px 0; padding: 8px 12px 8px 12px; background-color: gray; border-radius: 10px; color:white; ">{prompt}</div></div>',
|
||||
unsafe_allow_html=True)
|
||||
messages.append({"role": "user", "content": prompt})
|
||||
st.session_state.chat_messages.append({"role": "user", "content": prompt})
|
||||
messages.append({"role": "user", "content": prompt[-st.session_state.max_new_tokens:]})
|
||||
st.session_state.chat_messages.append({"role": "user", "content": prompt[-st.session_state.max_new_tokens:]})
|
||||
|
||||
with st.chat_message("assistant", avatar=image_url):
|
||||
placeholder = st.empty()
|
||||
random_seed = random.randint(0, 2 ** 32 - 1)
|
||||
setup_seed(random_seed)
|
||||
|
||||
st.session_state.chat_messages = system_prompt + st.session_state.chat_messages[
|
||||
-(st.session_state.history_chat_num + 1):]
|
||||
new_prompt = tokenizer.apply_chat_template(
|
||||
st.session_state.chat_messages,
|
||||
tokenize=False,
|
||||
add_generation_prompt=True
|
||||
)[-(st.session_state.max_new_tokens - 1):]
|
||||
|
||||
x = torch.tensor(tokenizer(new_prompt)['input_ids'], device=device).unsqueeze(0)
|
||||
with torch.no_grad():
|
||||
res_y = model.generate(x, tokenizer.eos_token_id, max_new_tokens=st.session_state.max_new_tokens,
|
||||
temperature=st.session_state.temperature,
|
||||
top_p=st.session_state.top_p, stream=True)
|
||||
if model_source == "API":
|
||||
try:
|
||||
for y in res_y:
|
||||
answer = tokenizer.decode(y[0].tolist(), skip_special_tokens=True)
|
||||
if (answer and answer[-1] == '�') or not answer:
|
||||
continue
|
||||
from openai import OpenAI
|
||||
|
||||
client = OpenAI(
|
||||
api_key=api_key,
|
||||
base_url=api_url
|
||||
)
|
||||
history_num = st.session_state.history_chat_num + 1 # +1 是为了包含当前的用户消息
|
||||
conversation_history = system_prompt + st.session_state.chat_messages[-history_num:]
|
||||
answer = ""
|
||||
response = client.chat.completions.create(
|
||||
model=api_model_id,
|
||||
messages=conversation_history,
|
||||
stream=True,
|
||||
temperature=st.session_state.temperature
|
||||
)
|
||||
|
||||
for chunk in response:
|
||||
content = chunk.choices[0].delta.content or ""
|
||||
answer += content
|
||||
placeholder.markdown(process_assistant_content(answer), unsafe_allow_html=True)
|
||||
except StopIteration:
|
||||
print("No answer")
|
||||
|
||||
assistant_answer = answer.replace(new_prompt, "")
|
||||
messages.append({"role": "assistant", "content": assistant_answer})
|
||||
st.session_state.chat_messages.append({"role": "assistant", "content": assistant_answer})
|
||||
except Exception as e:
|
||||
answer = f"API调用出错: {str(e)}"
|
||||
placeholder.markdown(answer, unsafe_allow_html=True)
|
||||
else:
|
||||
random_seed = random.randint(0, 2 ** 32 - 1)
|
||||
setup_seed(random_seed)
|
||||
|
||||
with st.empty():
|
||||
if st.button("×", key=f"delete_{len(messages) - 1}"):
|
||||
st.session_state.messages = st.session_state.messages[:-2]
|
||||
st.session_state.chat_messages = st.session_state.chat_messages[:-2]
|
||||
st.rerun()
|
||||
st.session_state.chat_messages = system_prompt + st.session_state.chat_messages[
|
||||
-(st.session_state.history_chat_num + 1):]
|
||||
new_prompt = tokenizer.apply_chat_template(
|
||||
st.session_state.chat_messages,
|
||||
tokenize=False,
|
||||
add_generation_prompt=True
|
||||
)
|
||||
|
||||
inputs = tokenizer(
|
||||
new_prompt,
|
||||
return_tensors="pt",
|
||||
truncation=True
|
||||
).to(device)
|
||||
|
||||
streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True)
|
||||
generation_kwargs = {
|
||||
"input_ids": inputs.input_ids,
|
||||
"max_length": inputs.input_ids.shape[1] + st.session_state.max_new_tokens,
|
||||
"num_return_sequences": 1,
|
||||
"do_sample": True,
|
||||
"attention_mask": inputs.attention_mask,
|
||||
"pad_token_id": tokenizer.pad_token_id,
|
||||
"eos_token_id": tokenizer.eos_token_id,
|
||||
"temperature": st.session_state.temperature,
|
||||
"top_p": 0.85,
|
||||
"streamer": streamer,
|
||||
}
|
||||
|
||||
Thread(target=model.generate, kwargs=generation_kwargs).start()
|
||||
|
||||
answer = ""
|
||||
for new_text in streamer:
|
||||
answer += new_text
|
||||
placeholder.markdown(process_assistant_content(answer), unsafe_allow_html=True)
|
||||
|
||||
messages.append({"role": "assistant", "content": answer})
|
||||
st.session_state.chat_messages.append({"role": "assistant", "content": answer})
|
||||
with st.empty():
|
||||
if st.button("×", key=f"delete_{len(messages) - 1}"):
|
||||
st.session_state.messages = st.session_state.messages[:-2]
|
||||
st.session_state.chat_messages = st.session_state.chat_messages[:-2]
|
||||
st.rerun()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer
|
||||
|
||||
main()
|
||||
|
||||
Reference in New Issue
Block a user