feat: 更新 README,添加生成长度控制和测试指令;修改 main_chat.py,优化模型加载和生成逻辑;新增单元测试文件 test_smoke.py

This commit is contained in:
2026-07-08 11:48:19 +08:00
parent 1de4656eb8
commit 24b56ebb6b
5 changed files with 177 additions and 20 deletions

View File

@@ -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

View File

@@ -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:

View File

@@ -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__':

View File

@@ -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
View 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()