Files
JspaceAI/main_chat.py
XiuchengWu 300e86956b Implement consensus and memory management for workspace events
- Add ConsensusSnapshot and ConsensusSlot classes for structured consensus state representation.
- Introduce OnlineLanguageLearner for stateful online learning with experience replay and EWC regularization.
- Create WorkspaceEvent and ActionEvent classes for serializable event types.
- Develop InMemoryVectorMemoryStore for storing and querying workspace events.
- Implement ActionPolicy to manage action selection and motor parameterization.
- Establish a unified WorkspaceRuntime for managing the agent's lifecycle and interactions.
- Enhance tests in test_smoke.py to cover new functionalities and ensure correctness.
2026-07-08 12:24:26 +08:00

276 lines
9.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/usr/bin/env python3
"""
JspaceAI —— 对话版(只有语言,控制台交互)
基于字符级语言模型 + 自主进化EWC + 经验回放 + 专家可塑性)。
交互式对话:用户输入文本,模型生成回复,同时持续学习。
模式:
--mode chat: 交互对话(默认)
--mode train: 先在 Shakespeare 语料上预训练若干步,再进入对话
--mode generate: 给定提示词一次性生成
运行:
python main_chat.py
python main_chat.py --mode train --steps 100
python main_chat.py --mode generate --prompt "To be"
"""
from __future__ import annotations
import argparse
import torch
from pathlib import Path
from jspaceai import (
LanguageConfig, JSpaceLanguageModel, EvolutionTrainer,
OnlineLanguageLearner,
CharTokenizer, load_chinese_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:
return LanguageConfig(
vocab_size=vocab_size,
embed_dim=64, input_dim=32,
workspace_dim=128, expert_dim=64,
num_experts=12, num_wells=8,
ode_steps=4, dt=0.1, tau_w=0.6,
jacobian_sparsity=32, noise_std=0.001,
use_rk4=True, use_layer_norm=True,
)
def load_or_init_model(device: str):
"""加载已保存的模型或初始化新模型"""
mp = Path('outputs/chat_model.pt')
if mp.exists():
try:
ckpt = torch.load(mp, map_location=device, weights_only=False)
# 用保存的 tokenizer chars 确保一致
saved_chars = ckpt.get('tokenizer_chars', None)
if 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)
model.load_state_dict(ckpt['model'])
print(f"已加载模型: {mp}vocab={tokenizer.vocab_size}")
return model, config, tokenizer, text
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)
print("全新初始化(首次对话)")
return model, config, tokenizer, text
def save_model(model, config, tokenizer):
Path('outputs').mkdir(exist_ok=True)
torch.save({
'model': model.state_dict(),
'config': config,
'tokenizer_chars': tokenizer.chars,
}, 'outputs/chat_model.pt')
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。
"""
print("\n" + "=" * 60)
print(f"预训练 {n_steps}Shakespeare 语料)")
print("=" * 60)
text = ensure_corpus(text)
config = model.config
# 训练时临时关 RK4 加速Euler 快 4 倍)
original_rk4 = config.use_rk4
for expert in model.experts:
expert.use_rk4 = False
print(f"训练模式: Euler加速训练后恢复 RK4")
trainer = EvolutionTrainer(
model, config, lr=5e-3, ewc_lambda=0.05, device=device,
)
chunks = [text[i:i+200] for i in range(0, len(text), 200)]
trainer.evolve(
chunks, tokenizer,
seq_len=64, batch_size=8,
consolidate_every=50, generate_every=50,
max_steps=n_steps, prompt_text="学而时习之",
)
# 恢复 RK4
for expert in model.experts:
expert.use_rk4 = original_rk4
save_model(model, config, tokenizer)
print(f"\n模型已保存: outputs/chat_model.pt")
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]
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 chat(model, config, tokenizer, text, device: str):
"""交互对话循环"""
print("\n" + "=" * 60)
print("JspaceAI 对话模式")
print("=" * 60)
print("输入文本与模型对话,模型会持续学习你的输入。")
print("命令: /quit 退出 /save 保存 /reset 重置 /train N 预训练N步")
print("=" * 60 + "\n")
model.eval()
learner = OnlineLanguageLearner(model, config, tokenizer, device=device)
while True:
try:
user_input = input("你: ").strip()
except (EOFError, KeyboardInterrupt):
print("\n再见")
break
if not user_input:
continue
if user_input.startswith('/'):
cmd = user_input.lower()
if cmd in ('/quit', '/exit', '/q'):
save_model(model, config, tokenizer)
print(f"(状态已保存到 outputs/chat_model.pt)")
print("再见")
break
elif cmd == '/save':
save_model(model, config, tokenizer)
print(f"(已保存到 outputs/chat_model.pt)")
continue
elif cmd == '/reset':
model = JSpaceLanguageModel(config).to(device)
learner = OnlineLanguageLearner(model, config, tokenizer, device=device)
print("(模型已重置为随机初始化)")
continue
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)
learner = OnlineLanguageLearner(model, config, tokenizer, device=device)
model.eval()
continue
else:
print("未知命令。可用: /quit /save /reset /train N")
continue
# 在线学习用户输入
learn_stats = learner.learn_text(user_input)
# 生成回复
response = generate_response(
model, tokenizer, user_input,
n_new=40, temperature=0.8, top_k=5,
)
print(f"AI: {response}")
if learn_stats is not None:
print(f" (学习 loss={learn_stats['loss']:.3f} replay={learn_stats['replay_loss']:.3f} step={learn_stats['step']})")
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, fast=fast)
print(f"提示: {prompt}")
print(f"生成: {response}")
def main():
p = argparse.ArgumentParser(description='JspaceAI 对话版')
p.add_argument('--mode', default='chat',
choices=['chat', 'train', 'generate'],
help='运行模式: chat=交互, train=预训练, generate=一次性生成')
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
if dev == 'auto':
dev = 'cuda' if torch.cuda.is_available() else (
'mps' if torch.backends.mps.is_available() else 'cpu')
model, config, tokenizer, text = load_or_init_model(dev)
if args.mode == 'chat':
chat(model, config, tokenizer, text, dev)
elif args.mode == 'train':
train(model, tokenizer, text, args.steps, dev)
elif args.mode == 'generate':
generate_once(model, tokenizer, args.prompt,
n_new=args.n_new, fast=not args.accurate)
if __name__ == '__main__':
main()