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"
|
||||
|
||||
# 控制生成长度;默认使用快速 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
|
||||
```
|
||||
|
||||
## 测试
|
||||
|
||||
```bash
|
||||
python -m unittest discover -s tests -v
|
||||
```
|
||||
|
||||
## 外挂模块系统(可热插拔)
|
||||
|
||||
```python
|
||||
|
||||
@@ -436,11 +436,13 @@ class MultimodalJSpaceModel(nn.Module):
|
||||
x = self.encode_modality(modality, data)
|
||||
|
||||
# 处理序列或单步
|
||||
w_traj_all = []
|
||||
if x.dim() == 2: # 单步 (batch, input_dim)
|
||||
state, w_traj = self.step(state, x, record_trajectory)
|
||||
if record_trajectory:
|
||||
w_traj_all = w_traj
|
||||
w = state['w']
|
||||
else: # 序列 (batch, T, input_dim)
|
||||
w_traj_all = []
|
||||
for t in range(x.shape[1]):
|
||||
state, w_traj = self.step(state, x[:, t], record_trajectory)
|
||||
if record_trajectory:
|
||||
|
||||
82
main_chat.py
82
main_chat.py
@@ -18,7 +18,6 @@ JspaceAI —— 对话版(只有语言,控制台交互)
|
||||
from __future__ import annotations
|
||||
import argparse
|
||||
import torch
|
||||
import numpy as np
|
||||
from pathlib import Path
|
||||
|
||||
from jspaceai import (
|
||||
@@ -28,6 +27,14 @@ from jspaceai import (
|
||||
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:
|
||||
return LanguageConfig(
|
||||
vocab_size=vocab_size,
|
||||
@@ -42,8 +49,6 @@ def get_config(vocab_size: int) -> LanguageConfig:
|
||||
|
||||
def load_or_init_model(device: str):
|
||||
"""加载已保存的模型或初始化新模型"""
|
||||
text = clean_corpus() # 清洗后语料(繁简统一+过滤)
|
||||
|
||||
mp = Path('outputs/chat_model.pt')
|
||||
if mp.exists():
|
||||
try:
|
||||
@@ -51,12 +56,11 @@ def load_or_init_model(device: str):
|
||||
# 用保存的 tokenizer chars 确保一致
|
||||
saved_chars = ckpt.get('tokenizer_chars', None)
|
||||
if saved_chars:
|
||||
tokenizer = CharTokenizer(
|
||||
chars=saved_chars,
|
||||
char_to_idx={c: i for i, c in enumerate(saved_chars)},
|
||||
idx_to_char={i: c for i, c in enumerate(saved_chars)},
|
||||
)
|
||||
tokenizer = tokenizer_from_chars(saved_chars)
|
||||
text = None
|
||||
else:
|
||||
print("checkpoint 缺少 tokenizer,正在准备语料...")
|
||||
text = clean_corpus()
|
||||
tokenizer = CharTokenizer.from_text(text)
|
||||
config = get_config(tokenizer.vocab_size)
|
||||
model = JSpaceLanguageModel(config).to(device)
|
||||
@@ -66,6 +70,7 @@ def load_or_init_model(device: str):
|
||||
except Exception as e:
|
||||
print(f"模型加载失败: {e},全新初始化")
|
||||
|
||||
text = clean_corpus() # 清洗后语料(繁简统一+过滤)
|
||||
tokenizer = CharTokenizer.from_text(text)
|
||||
config = get_config(tokenizer.vocab_size)
|
||||
model = JSpaceLanguageModel(config).to(device)
|
||||
@@ -82,7 +87,33 @@ def save_model(model, config, tokenizer):
|
||||
}, '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 语料上预训练
|
||||
|
||||
训练时临时关闭 RK4 用 Euler 加速(快 4 倍),训练完恢复 RK4。
|
||||
@@ -91,6 +122,7 @@ def train(model, tokenizer, text, n_steps: int, device: str):
|
||||
print(f"预训练 {n_steps} 步(Shakespeare 语料)")
|
||||
print("=" * 60)
|
||||
|
||||
text = ensure_corpus(text)
|
||||
config = model.config
|
||||
|
||||
# 训练时临时关 RK4 加速(Euler 快 4 倍)
|
||||
@@ -118,20 +150,25 @@ def train(model, tokenizer, text, n_steps: int, device: str):
|
||||
print(f"\n模型已保存: outputs/chat_model.pt")
|
||||
|
||||
|
||||
def generate_response(model, tokenizer, prompt: str, n_new: int = 100,
|
||||
temperature: float = 0.8, top_k: int = 5) -> str:
|
||||
def generate_response(model, tokenizer, prompt: str, n_new: int = 60,
|
||||
temperature: float = 0.8, top_k: int = 5,
|
||||
fast: bool = True) -> str:
|
||||
"""生成回复"""
|
||||
# 把用户输入编码(未知字符用 0)
|
||||
prompt_ids = tokenizer.encode(prompt)
|
||||
if not prompt_ids:
|
||||
prompt_ids = [0]
|
||||
generated = model.generate(
|
||||
prompt_ids, n_new=n_new, temperature=temperature, top_k=top_k,
|
||||
)
|
||||
was_training = model.training
|
||||
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)
|
||||
|
||||
|
||||
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 防遗忘"""
|
||||
import torch.nn.functional as F
|
||||
# 把用户输入作为新语料学习
|
||||
@@ -206,6 +243,7 @@ def chat(model, config, tokenizer, text, device: str):
|
||||
elif cmd.startswith('/train'):
|
||||
parts = cmd.split()
|
||||
n = int(parts[1]) if len(parts) > 1 else 50
|
||||
text = ensure_corpus(text)
|
||||
train(model, tokenizer, text, n, device)
|
||||
model.eval()
|
||||
continue
|
||||
@@ -219,18 +257,19 @@ def chat(model, config, tokenizer, text, device: str):
|
||||
# 生成回复
|
||||
response = generate_response(
|
||||
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}")
|
||||
if loss is not None:
|
||||
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()
|
||||
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"生成: {response}")
|
||||
|
||||
@@ -243,6 +282,10 @@ def main():
|
||||
p.add_argument('--steps', type=int, default=600, help='train 模式步数')
|
||||
p.add_argument('--prompt', default='To be', help='generate 模式提示词')
|
||||
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()
|
||||
|
||||
dev = args.device
|
||||
@@ -257,7 +300,8 @@ def main():
|
||||
elif args.mode == 'train':
|
||||
train(model, tokenizer, text, args.steps, dev)
|
||||
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__':
|
||||
|
||||
@@ -7,3 +7,4 @@ soundfile>=0.12
|
||||
mss>=9.0
|
||||
pynput>=1.7
|
||||
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