feat: 更新 README,添加生成长度控制和测试指令;修改 main_chat.py,优化模型加载和生成逻辑;新增单元测试文件 test_smoke.py
This commit is contained in:
10
README.md
10
README.md
@@ -123,6 +123,10 @@ python main_chat.py --mode chat
|
|||||||
|
|
||||||
# 一次性生成
|
# 一次性生成
|
||||||
python main_chat.py --mode generate --prompt "To be"
|
python main_chat.py --mode generate --prompt "To be"
|
||||||
|
|
||||||
|
# 控制生成长度;默认使用快速 Euler 推理,--accurate 改用 RK4
|
||||||
|
python main_chat.py --mode generate --prompt "学而时习之" --n-new 40
|
||||||
|
python main_chat.py --mode generate --prompt "学而时习之" --accurate
|
||||||
```
|
```
|
||||||
|
|
||||||
对话特性:
|
对话特性:
|
||||||
@@ -140,6 +144,12 @@ AI: ...
|
|||||||
你: /quit
|
你: /quit
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## 测试
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python -m unittest discover -s tests -v
|
||||||
|
```
|
||||||
|
|
||||||
## 外挂模块系统(可热插拔)
|
## 外挂模块系统(可热插拔)
|
||||||
|
|
||||||
```python
|
```python
|
||||||
|
|||||||
@@ -436,11 +436,13 @@ class MultimodalJSpaceModel(nn.Module):
|
|||||||
x = self.encode_modality(modality, data)
|
x = self.encode_modality(modality, data)
|
||||||
|
|
||||||
# 处理序列或单步
|
# 处理序列或单步
|
||||||
|
w_traj_all = []
|
||||||
if x.dim() == 2: # 单步 (batch, input_dim)
|
if x.dim() == 2: # 单步 (batch, input_dim)
|
||||||
state, w_traj = self.step(state, x, record_trajectory)
|
state, w_traj = self.step(state, x, record_trajectory)
|
||||||
|
if record_trajectory:
|
||||||
|
w_traj_all = w_traj
|
||||||
w = state['w']
|
w = state['w']
|
||||||
else: # 序列 (batch, T, input_dim)
|
else: # 序列 (batch, T, input_dim)
|
||||||
w_traj_all = []
|
|
||||||
for t in range(x.shape[1]):
|
for t in range(x.shape[1]):
|
||||||
state, w_traj = self.step(state, x[:, t], record_trajectory)
|
state, w_traj = self.step(state, x[:, t], record_trajectory)
|
||||||
if record_trajectory:
|
if record_trajectory:
|
||||||
|
|||||||
82
main_chat.py
82
main_chat.py
@@ -18,7 +18,6 @@ JspaceAI —— 对话版(只有语言,控制台交互)
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
import argparse
|
import argparse
|
||||||
import torch
|
import torch
|
||||||
import numpy as np
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from jspaceai import (
|
from jspaceai import (
|
||||||
@@ -28,6 +27,14 @@ from jspaceai import (
|
|||||||
from train_chat import clean_corpus
|
from train_chat import clean_corpus
|
||||||
|
|
||||||
|
|
||||||
|
def tokenizer_from_chars(chars: list[str]) -> CharTokenizer:
|
||||||
|
return CharTokenizer(
|
||||||
|
chars=chars,
|
||||||
|
char_to_idx={c: i for i, c in enumerate(chars)},
|
||||||
|
idx_to_char={i: c for i, c in enumerate(chars)},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def get_config(vocab_size: int) -> LanguageConfig:
|
def get_config(vocab_size: int) -> LanguageConfig:
|
||||||
return LanguageConfig(
|
return LanguageConfig(
|
||||||
vocab_size=vocab_size,
|
vocab_size=vocab_size,
|
||||||
@@ -42,8 +49,6 @@ def get_config(vocab_size: int) -> LanguageConfig:
|
|||||||
|
|
||||||
def load_or_init_model(device: str):
|
def load_or_init_model(device: str):
|
||||||
"""加载已保存的模型或初始化新模型"""
|
"""加载已保存的模型或初始化新模型"""
|
||||||
text = clean_corpus() # 清洗后语料(繁简统一+过滤)
|
|
||||||
|
|
||||||
mp = Path('outputs/chat_model.pt')
|
mp = Path('outputs/chat_model.pt')
|
||||||
if mp.exists():
|
if mp.exists():
|
||||||
try:
|
try:
|
||||||
@@ -51,12 +56,11 @@ def load_or_init_model(device: str):
|
|||||||
# 用保存的 tokenizer chars 确保一致
|
# 用保存的 tokenizer chars 确保一致
|
||||||
saved_chars = ckpt.get('tokenizer_chars', None)
|
saved_chars = ckpt.get('tokenizer_chars', None)
|
||||||
if saved_chars:
|
if saved_chars:
|
||||||
tokenizer = CharTokenizer(
|
tokenizer = tokenizer_from_chars(saved_chars)
|
||||||
chars=saved_chars,
|
text = None
|
||||||
char_to_idx={c: i for i, c in enumerate(saved_chars)},
|
|
||||||
idx_to_char={i: c for i, c in enumerate(saved_chars)},
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
|
print("checkpoint 缺少 tokenizer,正在准备语料...")
|
||||||
|
text = clean_corpus()
|
||||||
tokenizer = CharTokenizer.from_text(text)
|
tokenizer = CharTokenizer.from_text(text)
|
||||||
config = get_config(tokenizer.vocab_size)
|
config = get_config(tokenizer.vocab_size)
|
||||||
model = JSpaceLanguageModel(config).to(device)
|
model = JSpaceLanguageModel(config).to(device)
|
||||||
@@ -66,6 +70,7 @@ def load_or_init_model(device: str):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"模型加载失败: {e},全新初始化")
|
print(f"模型加载失败: {e},全新初始化")
|
||||||
|
|
||||||
|
text = clean_corpus() # 清洗后语料(繁简统一+过滤)
|
||||||
tokenizer = CharTokenizer.from_text(text)
|
tokenizer = CharTokenizer.from_text(text)
|
||||||
config = get_config(tokenizer.vocab_size)
|
config = get_config(tokenizer.vocab_size)
|
||||||
model = JSpaceLanguageModel(config).to(device)
|
model = JSpaceLanguageModel(config).to(device)
|
||||||
@@ -82,7 +87,33 @@ def save_model(model, config, tokenizer):
|
|||||||
}, 'outputs/chat_model.pt')
|
}, 'outputs/chat_model.pt')
|
||||||
|
|
||||||
|
|
||||||
def train(model, tokenizer, text, n_steps: int, device: str):
|
def ensure_corpus(text: str | None) -> str:
|
||||||
|
if text is None:
|
||||||
|
print("正在准备训练语料...")
|
||||||
|
return clean_corpus()
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
class temporary_rk4:
|
||||||
|
"""Temporarily switch expert integration mode for faster inference."""
|
||||||
|
|
||||||
|
def __init__(self, model, enabled: bool):
|
||||||
|
self.model = model
|
||||||
|
self.enabled = enabled
|
||||||
|
self.original = []
|
||||||
|
|
||||||
|
def __enter__(self):
|
||||||
|
self.original = [expert.use_rk4 for expert in self.model.experts]
|
||||||
|
for expert in self.model.experts:
|
||||||
|
expert.use_rk4 = self.enabled
|
||||||
|
|
||||||
|
def __exit__(self, exc_type, exc, tb):
|
||||||
|
for expert, enabled in zip(self.model.experts, self.original):
|
||||||
|
expert.use_rk4 = enabled
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def train(model, tokenizer, text: str | None, n_steps: int, device: str):
|
||||||
"""在 Shakespeare 语料上预训练
|
"""在 Shakespeare 语料上预训练
|
||||||
|
|
||||||
训练时临时关闭 RK4 用 Euler 加速(快 4 倍),训练完恢复 RK4。
|
训练时临时关闭 RK4 用 Euler 加速(快 4 倍),训练完恢复 RK4。
|
||||||
@@ -91,6 +122,7 @@ def train(model, tokenizer, text, n_steps: int, device: str):
|
|||||||
print(f"预训练 {n_steps} 步(Shakespeare 语料)")
|
print(f"预训练 {n_steps} 步(Shakespeare 语料)")
|
||||||
print("=" * 60)
|
print("=" * 60)
|
||||||
|
|
||||||
|
text = ensure_corpus(text)
|
||||||
config = model.config
|
config = model.config
|
||||||
|
|
||||||
# 训练时临时关 RK4 加速(Euler 快 4 倍)
|
# 训练时临时关 RK4 加速(Euler 快 4 倍)
|
||||||
@@ -118,20 +150,25 @@ def train(model, tokenizer, text, n_steps: int, device: str):
|
|||||||
print(f"\n模型已保存: outputs/chat_model.pt")
|
print(f"\n模型已保存: outputs/chat_model.pt")
|
||||||
|
|
||||||
|
|
||||||
def generate_response(model, tokenizer, prompt: str, n_new: int = 100,
|
def generate_response(model, tokenizer, prompt: str, n_new: int = 60,
|
||||||
temperature: float = 0.8, top_k: int = 5) -> str:
|
temperature: float = 0.8, top_k: int = 5,
|
||||||
|
fast: bool = True) -> str:
|
||||||
"""生成回复"""
|
"""生成回复"""
|
||||||
# 把用户输入编码(未知字符用 0)
|
# 把用户输入编码(未知字符用 0)
|
||||||
prompt_ids = tokenizer.encode(prompt)
|
prompt_ids = tokenizer.encode(prompt)
|
||||||
if not prompt_ids:
|
if not prompt_ids:
|
||||||
prompt_ids = [0]
|
prompt_ids = [0]
|
||||||
generated = model.generate(
|
was_training = model.training
|
||||||
prompt_ids, n_new=n_new, temperature=temperature, top_k=top_k,
|
with temporary_rk4(model, enabled=not fast):
|
||||||
)
|
generated = model.generate(
|
||||||
|
prompt_ids, n_new=n_new, temperature=temperature, top_k=top_k,
|
||||||
|
)
|
||||||
|
if not was_training:
|
||||||
|
model.eval()
|
||||||
return tokenizer.decode(generated)
|
return tokenizer.decode(generated)
|
||||||
|
|
||||||
|
|
||||||
def online_learn(model, config, tokenizer, text: str, user_input: str, device: str):
|
def online_learn(model, config, tokenizer, text: str | None, user_input: str, device: str):
|
||||||
"""在线学习用户输入 + 回放一段 Shakespeare 防遗忘"""
|
"""在线学习用户输入 + 回放一段 Shakespeare 防遗忘"""
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
# 把用户输入作为新语料学习
|
# 把用户输入作为新语料学习
|
||||||
@@ -206,6 +243,7 @@ def chat(model, config, tokenizer, text, device: str):
|
|||||||
elif cmd.startswith('/train'):
|
elif cmd.startswith('/train'):
|
||||||
parts = cmd.split()
|
parts = cmd.split()
|
||||||
n = int(parts[1]) if len(parts) > 1 else 50
|
n = int(parts[1]) if len(parts) > 1 else 50
|
||||||
|
text = ensure_corpus(text)
|
||||||
train(model, tokenizer, text, n, device)
|
train(model, tokenizer, text, n, device)
|
||||||
model.eval()
|
model.eval()
|
||||||
continue
|
continue
|
||||||
@@ -219,18 +257,19 @@ def chat(model, config, tokenizer, text, device: str):
|
|||||||
# 生成回复
|
# 生成回复
|
||||||
response = generate_response(
|
response = generate_response(
|
||||||
model, tokenizer, user_input,
|
model, tokenizer, user_input,
|
||||||
n_new=80, temperature=0.8, top_k=5,
|
n_new=40, temperature=0.8, top_k=5,
|
||||||
)
|
)
|
||||||
print(f"AI: {response}")
|
print(f"AI: {response}")
|
||||||
if loss is not None:
|
if loss is not None:
|
||||||
print(f" (学习 loss={loss:.3f})")
|
print(f" (学习 loss={loss:.3f})")
|
||||||
|
|
||||||
|
|
||||||
def generate_once(model, tokenizer, prompt: str, n_new: int = 200):
|
def generate_once(model, tokenizer, prompt: str, n_new: int = 80,
|
||||||
|
fast: bool = True):
|
||||||
"""一次性生成"""
|
"""一次性生成"""
|
||||||
model.eval()
|
model.eval()
|
||||||
response = generate_response(model, tokenizer, prompt, n_new=n_new,
|
response = generate_response(model, tokenizer, prompt, n_new=n_new,
|
||||||
temperature=0.7, top_k=5)
|
temperature=0.7, top_k=5, fast=fast)
|
||||||
print(f"提示: {prompt}")
|
print(f"提示: {prompt}")
|
||||||
print(f"生成: {response}")
|
print(f"生成: {response}")
|
||||||
|
|
||||||
@@ -243,6 +282,10 @@ def main():
|
|||||||
p.add_argument('--steps', type=int, default=600, help='train 模式步数')
|
p.add_argument('--steps', type=int, default=600, help='train 模式步数')
|
||||||
p.add_argument('--prompt', default='To be', help='generate 模式提示词')
|
p.add_argument('--prompt', default='To be', help='generate 模式提示词')
|
||||||
p.add_argument('--device', default='cpu', help='设备 (cpu/cuda/mps/auto)')
|
p.add_argument('--device', default='cpu', help='设备 (cpu/cuda/mps/auto)')
|
||||||
|
p.add_argument('--n-new', type=int, default=80,
|
||||||
|
help='generate 模式生成的新字符数')
|
||||||
|
p.add_argument('--accurate', action='store_true',
|
||||||
|
help='生成时使用 RK4(更慢但与训练配置一致)')
|
||||||
args = p.parse_args()
|
args = p.parse_args()
|
||||||
|
|
||||||
dev = args.device
|
dev = args.device
|
||||||
@@ -257,7 +300,8 @@ def main():
|
|||||||
elif args.mode == 'train':
|
elif args.mode == 'train':
|
||||||
train(model, tokenizer, text, args.steps, dev)
|
train(model, tokenizer, text, args.steps, dev)
|
||||||
elif args.mode == 'generate':
|
elif args.mode == 'generate':
|
||||||
generate_once(model, tokenizer, args.prompt)
|
generate_once(model, tokenizer, args.prompt,
|
||||||
|
n_new=args.n_new, fast=not args.accurate)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
|
|||||||
@@ -7,3 +7,4 @@ soundfile>=0.12
|
|||||||
mss>=9.0
|
mss>=9.0
|
||||||
pynput>=1.7
|
pynput>=1.7
|
||||||
av>=10.0
|
av>=10.0
|
||||||
|
opencc-python-reimplemented>=0.1.7
|
||||||
|
|||||||
100
tests/test_smoke.py
Normal file
100
tests/test_smoke.py
Normal file
@@ -0,0 +1,100 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from jspaceai import (
|
||||||
|
CharTokenizer,
|
||||||
|
JSpaceConfig,
|
||||||
|
JSpaceLanguageModel,
|
||||||
|
JSpaceModel,
|
||||||
|
LanguageConfig,
|
||||||
|
MultimodalConfig,
|
||||||
|
MultimodalJSpaceModel,
|
||||||
|
)
|
||||||
|
from main_chat import generate_response
|
||||||
|
|
||||||
|
|
||||||
|
class SmokeTests(unittest.TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
torch.manual_seed(0)
|
||||||
|
|
||||||
|
def test_core_forward_shapes(self):
|
||||||
|
config = JSpaceConfig(
|
||||||
|
input_dim=8,
|
||||||
|
workspace_dim=16,
|
||||||
|
expert_dim=8,
|
||||||
|
num_experts=3,
|
||||||
|
ode_steps=2,
|
||||||
|
noise_std=0.0,
|
||||||
|
)
|
||||||
|
model = JSpaceModel(config)
|
||||||
|
xs = torch.randn(2, 5, 8)
|
||||||
|
|
||||||
|
preds, info = model(xs, record_trajectory=True)
|
||||||
|
|
||||||
|
self.assertEqual(tuple(preds.shape), (2, 5, 8))
|
||||||
|
self.assertEqual(tuple(info["alpha"].shape), (2, 5, 3))
|
||||||
|
self.assertEqual(tuple(info["w_trajectory"].shape), (2, 5, 2, 16))
|
||||||
|
|
||||||
|
def test_language_fast_generation_preserves_eval_and_rk4(self):
|
||||||
|
tokenizer = CharTokenizer.from_text("学而时习之")
|
||||||
|
config = LanguageConfig(
|
||||||
|
vocab_size=tokenizer.vocab_size,
|
||||||
|
embed_dim=8,
|
||||||
|
input_dim=8,
|
||||||
|
workspace_dim=16,
|
||||||
|
expert_dim=8,
|
||||||
|
num_experts=3,
|
||||||
|
ode_steps=1,
|
||||||
|
noise_std=0.0,
|
||||||
|
)
|
||||||
|
model = JSpaceLanguageModel(config)
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
response = generate_response(
|
||||||
|
model,
|
||||||
|
tokenizer,
|
||||||
|
"学",
|
||||||
|
n_new=3,
|
||||||
|
temperature=1.0,
|
||||||
|
top_k=2,
|
||||||
|
fast=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(len(response), 3)
|
||||||
|
self.assertFalse(model.training)
|
||||||
|
self.assertTrue(all(expert.use_rk4 for expert in model.experts))
|
||||||
|
|
||||||
|
def test_multimodal_single_step_records_trajectory(self):
|
||||||
|
config = MultimodalConfig(
|
||||||
|
vocab_size=20,
|
||||||
|
embed_dim=8,
|
||||||
|
input_dim=8,
|
||||||
|
workspace_dim=16,
|
||||||
|
expert_dim=8,
|
||||||
|
num_experts=4,
|
||||||
|
ode_steps=2,
|
||||||
|
noise_std=0.0,
|
||||||
|
audio_frame_size=128,
|
||||||
|
)
|
||||||
|
model = MultimodalJSpaceModel(config)
|
||||||
|
token = torch.tensor([1])
|
||||||
|
|
||||||
|
outputs, info = model.forward_multimodal(
|
||||||
|
"text",
|
||||||
|
token,
|
||||||
|
record_trajectory=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(tuple(outputs["w"].shape), (1, 16))
|
||||||
|
self.assertEqual(len(info["w_trajectory"]), 2)
|
||||||
|
|
||||||
|
def test_tokenizer_unknown_maps_to_zero(self):
|
||||||
|
tokenizer = CharTokenizer.from_text("abc")
|
||||||
|
|
||||||
|
self.assertEqual(tokenizer.encode("?"), [0])
|
||||||
|
self.assertEqual(tokenizer.decode([0]), "<unk>")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user