mirror of
https://github.com/jingyaogong/minimind.git
synced 2026-10-03 08:07:21 +00:00
[fix] robustness
This commit is contained in:
@@ -765,16 +765,16 @@ LoRA 是一种常见的参数高效微调(Parameter-Efficient Fine-Tuning, PEF
|
||||
```bash
|
||||
# train_lora.py 在 CPU 上通常也能比较轻快地完成
|
||||
# 方式1
|
||||
torchrun --nproc_per_node 1 train_lora.py
|
||||
cd trainer && torchrun --nproc_per_node 1 train_lora.py
|
||||
# 方式2
|
||||
python train_lora.py
|
||||
cd trainer && python train_lora.py
|
||||
```
|
||||
|
||||
> 训练后的模型权重文件默认每隔`save_interval步`保存为: `lora_xxx_*.pth`(*为模型具体dimension,每次保存时新文件会覆盖旧文件)
|
||||
|
||||
|
||||
LoRA 很适合处理“如何在尽量保留通用能力的前提下,让模型快速适应私有领域或垂直场景”这类问题。例如基础模型医学知识不足时,就可以在原有模型之上叠加一层面向医疗场景的 LoRA 权重,以较小代价获得更好的领域表现。
|
||||
通常只需要准备同样的多轮对话格式数据,放置到 `lora_xxx.jsonl`,再执行 `python train_lora.py`,即可得到新的 `LoRA` 模型权重。
|
||||
通常只需要准备同样的多轮对话格式数据,放置到 `lora_xxx.jsonl`,再从仓库根目录执行 `cd trainer && python train_lora.py`,即可得到新的 `LoRA` 模型权重。
|
||||
|
||||
例1:垂域数据
|
||||
|
||||
|
||||
+4
-4
@@ -765,16 +765,16 @@ Its core idea is to introduce low-rank incremental branches alongside the origin
|
||||
```bash
|
||||
# train_lora.py can usually be completed fairly quickly even on CPU
|
||||
# Method 1
|
||||
torchrun --nproc_per_node 1 train_lora.py
|
||||
cd trainer && torchrun --nproc_per_node 1 train_lora.py
|
||||
# Method 2
|
||||
python train_lora.py
|
||||
cd trainer && python train_lora.py
|
||||
```
|
||||
|
||||
> The trained model weight files are saved by default every `save_interval steps` as: `lora_xxx_*.pth` (* is the specific model dimension, each save overwrites the previous file)
|
||||
|
||||
|
||||
LoRA is well-suited for handling problems like "how to let the model quickly adapt to private domains or vertical scenarios while preserving general capabilities as much as possible." For example, when the base model lacks medical knowledge, a medical-oriented LoRA weight layer can be stacked on top of the original model to achieve better domain performance at relatively small cost.
|
||||
Usually you only need to prepare multi-turn dialogue format data in the same way, place it in `lora_xxx.jsonl`, and then run `python train_lora.py` to obtain new `LoRA` model weights.
|
||||
Usually you only need to prepare multi-turn dialogue format data in the same way, place it in `lora_xxx.jsonl`, and then run `cd trainer && python train_lora.py` from the repository root to obtain new `LoRA` model weights.
|
||||
|
||||
Example 1: Vertical domain data
|
||||
|
||||
@@ -1989,4 +1989,4 @@ If `MiniMind` has been helpful to your research or work, feel free to cite:
|
||||
|
||||
# ⚖️ License
|
||||
|
||||
This project is open-sourced under the [Apache License 2.0](LICENSE).
|
||||
This project is open-sourced under the [Apache License 2.0](LICENSE).
|
||||
|
||||
+7
-1
@@ -21,13 +21,19 @@ while True:
|
||||
extra_body={"chat_template_kwargs": {"open_thinking": True}, "reasoning_effort": "medium"} # 思考开关
|
||||
)
|
||||
if not stream:
|
||||
if not response.choices or response.choices[0].message is None:
|
||||
raise ValueError("LLM returned empty or filtered response")
|
||||
assistant_res = response.choices[0].message.content
|
||||
print('[A]: ', assistant_res)
|
||||
else:
|
||||
print('[A]: ', end='', flush=True)
|
||||
assistant_res = ''
|
||||
for chunk in response:
|
||||
if not chunk.choices:
|
||||
continue
|
||||
delta = chunk.choices[0].delta
|
||||
if delta is None:
|
||||
continue
|
||||
r = getattr(delta, 'reasoning_content', None) or ""
|
||||
c = delta.content or ""
|
||||
if r:
|
||||
@@ -37,4 +43,4 @@ while True:
|
||||
assistant_res += c
|
||||
|
||||
conversation_history.append({"role": "assistant", "content": assistant_res})
|
||||
print('\n\n')
|
||||
print('\n\n')
|
||||
|
||||
+22
-15
@@ -15,7 +15,7 @@ from threading import Thread
|
||||
from queue import Queue
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, Field
|
||||
from transformers import AutoTokenizer, AutoModelForCausalLM, TextStreamer
|
||||
from model.model_minimind import MiniMindConfig, MiniMindForCausalLM
|
||||
from model.model_lora import apply_lora, load_lora
|
||||
@@ -54,7 +54,7 @@ class ChatRequest(BaseModel):
|
||||
top_p: float = 0.92
|
||||
max_tokens: int = 8192
|
||||
stream: bool = True
|
||||
tools: list = []
|
||||
tools: list = Field(default_factory=list)
|
||||
open_thinking: bool = False
|
||||
chat_template_kwargs: dict = None
|
||||
|
||||
@@ -104,24 +104,28 @@ def parse_response(text):
|
||||
|
||||
def generate_stream_response(messages, temperature, top_p, max_tokens, tools=None, open_thinking=False):
|
||||
try:
|
||||
new_prompt = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True, tools=tools or None, open_thinking=open_thinking)[-max_tokens:]
|
||||
new_prompt = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True, tools=tools or None, open_thinking=open_thinking)
|
||||
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,
|
||||
attention_mask=inputs.attention_mask,
|
||||
pad_token_id=tokenizer.pad_token_id,
|
||||
eos_token_id=tokenizer.eos_token_id,
|
||||
streamer=streamer
|
||||
)
|
||||
try:
|
||||
model.generate(
|
||||
inputs.input_ids,
|
||||
max_new_tokens=max_tokens,
|
||||
do_sample=True,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
attention_mask=inputs.attention_mask,
|
||||
pad_token_id=tokenizer.pad_token_id,
|
||||
eos_token_id=tokenizer.eos_token_id,
|
||||
streamer=streamer
|
||||
)
|
||||
except Exception as e:
|
||||
queue.put({"error": str(e)})
|
||||
queue.put(None)
|
||||
|
||||
Thread(target=_generate).start()
|
||||
|
||||
@@ -133,6 +137,9 @@ def generate_stream_response(messages, temperature, top_p, max_tokens, tools=Non
|
||||
text = queue.get()
|
||||
if text is None:
|
||||
break
|
||||
if isinstance(text, dict):
|
||||
yield json.dumps(text, ensure_ascii=False)
|
||||
continue
|
||||
full_text += text
|
||||
|
||||
if not thinking_ended:
|
||||
@@ -190,7 +197,7 @@ async def chat_completions(request: ChatRequest):
|
||||
add_generation_prompt=True,
|
||||
tools=request.tools or None,
|
||||
open_thinking=request.get_open_thinking()
|
||||
)[-request.max_tokens:]
|
||||
)
|
||||
inputs = tokenizer(new_prompt, return_tensors="pt", truncation=True).to(device)
|
||||
with torch.no_grad():
|
||||
generated_ids = model.generate(
|
||||
|
||||
+2
-2
@@ -339,8 +339,8 @@ 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: #3d4450; border-radius: 22px; color: white;">{prompt}</div></div>',
|
||||
unsafe_allow_html=True)
|
||||
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:]})
|
||||
messages.append({"role": "user", "content": prompt})
|
||||
st.session_state.chat_messages.append({"role": "user", "content": prompt})
|
||||
|
||||
placeholder = st.empty()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user