This commit is contained in:
jingyaogong
2025-04-26 10:05:47 +08:00
parent 7da201a944
commit a62faf34bd
23 changed files with 1110 additions and 19633 deletions
+2 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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()