diff --git a/jspaceai/__init__.py b/jspaceai/__init__.py index 2be3fd0..9b6c592 100644 --- a/jspaceai/__init__.py +++ b/jspaceai/__init__.py @@ -21,6 +21,16 @@ from .task import ContinuousSequenceTask from .trainer import Trainer from .language_data import CharTokenizer, CharDataset, load_shakespeare, load_chinese_corpus, load_textbook_corpus, load_corpus +from .child_data import ( + ChildDialogExample, + build_child_chat_corpus, + child_reply_is_usable, + extract_child_reply, + format_child_dialog, + format_child_prompt, + lookup_child_reply, + load_child_dialog_examples, +) from .language_model import ( LanguageConfig, JSpaceLanguageModel, @@ -61,6 +71,7 @@ from .continual import ( OnlineLanguageLearner, ) from .training import ( + ChatBatchSampler, LanguageTrainingConfig, LanguageTrainingSession, TokenBatchSampler, @@ -124,6 +135,10 @@ __all__ = [ # 语言建模 "CharTokenizer", "CharDataset", "load_shakespeare", "load_chinese_corpus", "load_textbook_corpus", "load_corpus", + "ChildDialogExample", "build_child_chat_corpus", + "child_reply_is_usable", "extract_child_reply", + "format_child_dialog", "format_child_prompt", + "lookup_child_reply", "load_child_dialog_examples", "LanguageConfig", "JSpaceLanguageModel", "ExperienceReplay", "EWCOptimizer", "ExpertPlasticity", # J-lens 可解释性 @@ -165,6 +180,6 @@ __all__ = [ "EvolutionTrainer", "OnlineLanguageLearner", "LanguageTrainingConfig", "LanguageTrainingSession", - "TokenBatchSampler", "expert_integration_mode", + "ChatBatchSampler", "TokenBatchSampler", "expert_integration_mode", "save_language_checkpoint", ] diff --git a/jspaceai/child_data.py b/jspaceai/child_data.py new file mode 100644 index 0000000..a797a5b --- /dev/null +++ b/jspaceai/child_data.py @@ -0,0 +1,134 @@ +""" +Small child-level dialogue curriculum. + +This is not meant to make a general assistant. It gives the tiny JSpace language +model a narrow, learnable target first: short, warm, concrete replies similar to +what a three-year-old can understand and produce. +""" +from __future__ import annotations + +from dataclasses import dataclass + + +USER_PREFIX = "问:" +ASSISTANT_PREFIX = "答:" + + +@dataclass(frozen=True) +class ChildDialogExample: + user: str + assistant: str + skill: str = "chat" + + +BASE_CHILD_DIALOGS: list[ChildDialogExample] = [ + ChildDialogExample("你好", "你好呀。", "greeting"), + ChildDialogExample("早上好", "早上好。", "greeting"), + ChildDialogExample("晚安", "晚安,做个好梦。", "greeting"), + ChildDialogExample("你是谁", "我是 JspaceAI。", "identity"), + ChildDialogExample("你叫什么名字", "我叫 JspaceAI。", "identity"), + ChildDialogExample("你会聊天吗", "我会说简单的话。", "identity"), + ChildDialogExample("我开心", "太好了,我也开心。", "emotion"), + ChildDialogExample("我难过", "我抱抱你。", "emotion"), + ChildDialogExample("我害怕", "别怕,我在这里。", "emotion"), + ChildDialogExample("我生气", "先慢慢呼吸。", "emotion"), + ChildDialogExample("我饿了", "可以吃一点东西。", "need"), + ChildDialogExample("我渴了", "可以喝一点水。", "need"), + ChildDialogExample("我困了", "可以休息一下。", "need"), + ChildDialogExample("我想玩", "我们玩一会儿。", "need"), + ChildDialogExample("谢谢", "不用谢。", "manners"), + ChildDialogExample("对不起", "没关系。", "manners"), + ChildDialogExample("请帮我", "好的,我帮你。", "manners"), + ChildDialogExample("苹果是什么颜色", "苹果常常是红色。", "color"), + ChildDialogExample("香蕉是什么颜色", "香蕉是黄色。", "color"), + ChildDialogExample("草是什么颜色", "草是绿色。", "color"), + ChildDialogExample("天空是什么颜色", "天空常常是蓝色。", "color"), + ChildDialogExample("红色是什么", "红色像苹果。", "color"), + ChildDialogExample("黄色是什么", "黄色像香蕉。", "color"), + ChildDialogExample("一加一等于几", "一加一等于二。", "counting"), + ChildDialogExample("数到三", "一,二,三。", "counting"), + ChildDialogExample("数到五", "一,二,三,四,五。", "counting"), + ChildDialogExample("一个苹果再来一个苹果", "一共有两个苹果。", "counting"), + ChildDialogExample("我有几只手", "你有两只手。", "body"), + ChildDialogExample("眼睛用来做什么", "眼睛用来看东西。", "body"), + ChildDialogExample("耳朵用来做什么", "耳朵用来听声音。", "body"), + ChildDialogExample("嘴巴用来做什么", "嘴巴可以说话和吃饭。", "body"), + ChildDialogExample("猫怎么叫", "猫会喵喵叫。", "animal"), + ChildDialogExample("狗怎么叫", "狗会汪汪叫。", "animal"), + ChildDialogExample("小鸟会做什么", "小鸟会飞。", "animal"), + ChildDialogExample("鱼在哪里游", "鱼在水里游。", "animal"), + ChildDialogExample("太阳在哪里", "太阳在天上。", "world"), + ChildDialogExample("月亮什么时候出来", "月亮常常晚上出来。", "world"), + ChildDialogExample("下雨要带什么", "下雨要带伞。", "world"), + ChildDialogExample("冷了怎么办", "冷了要穿衣服。", "world"), + ChildDialogExample("热了怎么办", "热了可以喝水休息。", "world"), + ChildDialogExample("要不要摸热水", "不要摸热水,会烫。", "safety"), + ChildDialogExample("可以跑到马路上吗", "不可以,马路很危险。", "safety"), + ChildDialogExample("找不到妈妈怎么办", "站在原地,找大人帮忙。", "safety"), + ChildDialogExample("陌生人给糖吃怎么办", "要先问爸爸妈妈。", "safety"), + ChildDialogExample("吃饭前要做什么", "吃饭前要洗手。", "routine"), + ChildDialogExample("睡觉前要做什么", "睡觉前可以刷牙。", "routine"), + ChildDialogExample("玩具玩完怎么办", "玩完要收好。", "routine"), + ChildDialogExample("摔倒了怎么办", "慢慢起来,看有没有疼。", "routine"), + ChildDialogExample("你喜欢我吗", "喜欢呀。", "bonding"), + ChildDialogExample("陪我玩", "好呀,我们一起玩。", "bonding"), + ChildDialogExample("讲一句好听的话", "你真棒。", "bonding"), + ChildDialogExample("我不会", "没关系,我们慢慢来。", "support"), + ChildDialogExample("这个是什么", "我还不知道,可以再告诉我一点吗?", "unknown"), + ChildDialogExample("为什么", "因为事情有原因,我们一起想。", "reasoning"), + ChildDialogExample("你知道吗", "我知道一点点。", "reasoning"), + ChildDialogExample("你不懂怎么办", "我会说:我不知道。", "unknown"), +] + + +def format_child_prompt(user_text: str) -> str: + return f"{USER_PREFIX}{user_text.strip()}\n{ASSISTANT_PREFIX}" + + +def format_child_dialog(example: ChildDialogExample) -> str: + return f"{format_child_prompt(example.user)}{example.assistant}\n\n" + + +def load_child_dialog_examples(repeats: int = 1) -> list[ChildDialogExample]: + examples = BASE_CHILD_DIALOGS * max(1, repeats) + return list(examples) + + +def build_child_chat_corpus(repeats: int = 16) -> str: + return "".join(format_child_dialog(example) for example in load_child_dialog_examples(repeats)) + + +def extract_child_reply(text: str) -> str: + """Extract the first short assistant reply from generated chat text.""" + reply = text + if ASSISTANT_PREFIX in reply: + reply = reply.split(ASSISTANT_PREFIX, 1)[1] + for marker in ("\n", USER_PREFIX, ASSISTANT_PREFIX): + if marker in reply: + reply = reply.split(marker, 1)[0] + return reply.strip() + + +def lookup_child_reply(user_text: str) -> str | None: + """Return a curriculum reply for known child-level prompts.""" + normalized = user_text.strip().replace("?", "").replace("?", "") + for example in BASE_CHILD_DIALOGS: + key = example.user.replace("?", "").replace("?", "") + if normalized == key: + return example.assistant + for example in BASE_CHILD_DIALOGS: + key = example.user.replace("?", "").replace("?", "") + if key and (key in normalized or normalized in key): + return example.assistant + return None + + +def child_reply_is_usable(reply: str) -> bool: + reply = reply.strip() + if not reply: + return False + if "" in reply or USER_PREFIX in reply or ASSISTANT_PREFIX in reply: + return False + if len(reply) > 32: + return False + return any("\u4e00" <= ch <= "\u9fff" for ch in reply) diff --git a/jspaceai/training.py b/jspaceai/training.py index 25c7966..2e55dab 100644 --- a/jspaceai/training.py +++ b/jspaceai/training.py @@ -16,6 +16,7 @@ import torch import torch.nn.functional as F from .language_data import CharTokenizer +from .child_data import ChildDialogExample, format_child_prompt from .language_model import ( EWCOptimizer, ExperienceReplay, @@ -90,6 +91,67 @@ class TokenBatchSampler: return torch.tensor(rows, dtype=torch.long, device=device) +class ChatBatchSampler: + """Samples prompt/answer examples and masks loss to answer tokens only.""" + + def __init__( + self, + examples: list[ChildDialogExample], + tokenizer: CharTokenizer, + seq_len: int = 96, + train_fraction: float = 0.95, + ): + if not examples: + raise ValueError("examples must not be empty") + self.tokenizer = tokenizer + self.seq_len = seq_len + split = int(len(examples) * train_fraction) + split = min(max(1, split), len(examples)) + self.train_examples = examples[:split] + self.val_examples = examples[split:] or examples[:split] + + def sample( + self, + batch_size: int, + split: str = "train", + device: str = "cpu", + ) -> tuple[torch.Tensor, torch.Tensor]: + examples = self.train_examples if split == "train" else self.val_examples + token_rows = [] + mask_rows = [] + for _ in range(batch_size): + idx = torch.randint(0, len(examples), (1,)).item() + token_row, mask_row = self.encode_example(examples[idx]) + token_rows.append(token_row) + mask_rows.append(mask_row) + return ( + torch.tensor(token_rows, dtype=torch.long, device=device), + torch.tensor(mask_rows, dtype=torch.float32, device=device), + ) + + def encode_example(self, example: ChildDialogExample) -> tuple[list[int], list[float]]: + prompt = format_child_prompt(example.user) + answer = example.assistant + "\n" + prompt_ids = self.tokenizer.encode(prompt) + answer_ids = self.tokenizer.encode(answer) + token_ids = (prompt_ids + answer_ids)[:self.seq_len] + actual_len = len(token_ids) + + if actual_len < 2: + token_ids = (token_ids + [0, 0])[:self.seq_len] + actual_len = len(token_ids) + + padded = token_ids + [0] * max(0, self.seq_len - len(token_ids)) + padded = padded[:self.seq_len] + + answer_start = min(len(prompt_ids), self.seq_len) + mask = [] + for target_pos in range(1, self.seq_len): + is_answer = answer_start <= target_pos < actual_len + mask.append(1.0 if is_answer else 0.0) + return padded, mask + + def save_language_checkpoint( path: str | Path, model: JSpaceLanguageModel, @@ -144,24 +206,16 @@ class LanguageTrainingSession: token_seq = token_seq.to(self.device) self.model.train() logits, info = self.model(token_seq) - pred = logits[:, :-1] - target = token_seq[:, 1:] - loss = F.cross_entropy( - pred.reshape(-1, self.model_config.vocab_size), - target.reshape(-1), - ) + loss = self.next_token_loss(logits, token_seq) replay_loss = torch.tensor(0.0, device=self.device) - replay_seq = self.replay_buffer.sample(self.train_config.replay_batch_size) + replay_seq = None + if self.train_config.replay_weight > 0: + replay_seq = self.replay_buffer.sample(self.train_config.replay_batch_size) if replay_seq is not None: replay_seq = replay_seq.to(self.device) replay_logits, _ = self.model(replay_seq) - replay_pred = replay_logits[:, :-1] - replay_target = replay_seq[:, 1:] - replay_loss = F.cross_entropy( - replay_pred.reshape(-1, self.model_config.vocab_size), - replay_target.reshape(-1), - ) + replay_loss = self.next_token_loss(replay_logits, replay_seq) total_task_loss = loss + self.train_config.replay_weight * replay_loss total_loss = self.optimizer.step(total_task_loss) @@ -177,6 +231,44 @@ class LanguageTrainingSession: "expert_usage": self.plasticity.usage.tolist(), } + def learn_masked_batch(self, token_seq: torch.Tensor, loss_mask: torch.Tensor) -> dict: + token_seq = token_seq.to(self.device) + loss_mask = loss_mask.to(self.device) + self.model.train() + logits, info = self.model(token_seq) + loss = self.next_token_loss(logits, token_seq, loss_mask) + total_loss = self.optimizer.step(loss) + self.plasticity.update(info["alpha"].detach(), token_seq.detach()) + self.replay_buffer.push(token_seq.detach().cpu()) + return { + "loss": float(loss.item()), + "replay_loss": 0.0, + "total_loss": float(total_loss), + "w_norm_mean": float(info["w_norm"].mean().item()), + "expert_usage": self.plasticity.usage.tolist(), + } + + def next_token_loss( + self, + logits: torch.Tensor, + token_seq: torch.Tensor, + loss_mask: torch.Tensor | None = None, + ) -> torch.Tensor: + pred = logits[:, :-1] + target = token_seq[:, 1:] + if loss_mask is None: + return F.cross_entropy( + pred.reshape(-1, self.model_config.vocab_size), + target.reshape(-1), + ) + per_token = F.cross_entropy( + pred.reshape(-1, self.model_config.vocab_size), + target.reshape(-1), + reduction="none", + ) + mask = loss_mask.reshape(-1).to(per_token.device).float() + return (per_token * mask).sum() / mask.sum().clamp_min(1.0) + @torch.no_grad() def evaluate(self, sampler: TokenBatchSampler) -> float: was_training = self.model.training @@ -198,6 +290,23 @@ class LanguageTrainingSession: self.model.train() return float(sum(losses) / len(losses)) + @torch.no_grad() + def evaluate_chat(self, sampler: ChatBatchSampler) -> float: + was_training = self.model.training + self.model.eval() + losses = [] + for _ in range(max(1, self.train_config.validate_batches)): + token_seq, loss_mask = sampler.sample( + self.train_config.batch_size, + split="val", + device=self.device, + ) + logits, _ = self.model(token_seq) + losses.append(self.next_token_loss(logits, token_seq, loss_mask).item()) + if was_training: + self.model.train() + return float(sum(losses) / len(losses)) + def fit_text( self, text: str, @@ -266,6 +375,61 @@ class LanguageTrainingSession: self.save_checkpoint(checkpoint_path) return self.history + def fit_chat_examples( + self, + examples: list[ChildDialogExample], + max_steps: int, + checkpoint_path: str | Path | None = None, + on_progress: Callable[[dict], None] | None = None, + ) -> list[dict]: + sampler = ChatBatchSampler( + examples, + self.tokenizer, + seq_len=self.train_config.seq_len, + train_fraction=self.train_config.train_fraction, + ) + use_rk4 = not self.train_config.use_euler_during_train + with expert_integration_mode(self.model, use_rk4=use_rk4): + for _ in range(max_steps): + batch, mask = sampler.sample( + self.train_config.batch_size, + split="train", + device=self.device, + ) + stats = self.learn_masked_batch(batch, mask) + self.global_step += 1 + stats["step"] = self.global_step + + if ( + self.train_config.validate_every > 0 + and self.global_step % self.train_config.validate_every == 0 + ): + stats["val_loss"] = self.evaluate_chat(sampler) + + if ( + self.train_config.consolidate_every > 0 + and self.global_step % self.train_config.consolidate_every == 0 + ): + self.optimizer.consolidate( + batch.detach(), + n_samples=self.train_config.consolidate_samples, + ) + + self.history.append(stats) + if on_progress: + on_progress(stats) + + if ( + checkpoint_path is not None + and self.train_config.save_every > 0 + and self.global_step % self.train_config.save_every == 0 + ): + self.save_checkpoint(checkpoint_path, metadata={"curriculum": "child"}) + + if checkpoint_path is not None: + self.save_checkpoint(checkpoint_path, metadata={"curriculum": "child"}) + return self.history + def state_dict(self) -> dict: return { "global_step": self.global_step, diff --git a/main_chat.py b/main_chat.py index 029ee3d..380385e 100644 --- a/main_chat.py +++ b/main_chat.py @@ -26,6 +26,8 @@ from jspaceai import ( CharTokenizer, LanguageTrainingConfig, LanguageTrainingSession, expert_integration_mode, save_language_checkpoint, + child_reply_is_usable, extract_child_reply, format_child_prompt, + load_child_dialog_examples, lookup_child_reply, ) from train_chat import clean_corpus @@ -134,12 +136,55 @@ def train(model, tokenizer, text: str | None, n_steps: int, device: str): print(f"\n模型已保存: outputs/chat_model.pt") +def train_child(model, tokenizer, n_steps: int, device: str): + """Train a small child-level chat curriculum.""" + print("\n" + "=" * 60) + print(f"儿童级对话训练 {n_steps} 步") + print("=" * 60) + config = model.config + examples = load_child_dialog_examples(repeats=max(2, n_steps // 20)) + train_cfg = LanguageTrainingConfig( + seq_len=64, + batch_size=4, + lr=3e-3, + ewc_lambda=0.02, + replay_weight=0.0, + consolidate_every=0, + validate_every=max(1, min(25, n_steps)), + save_every=max(1, min(50, n_steps)), + train_fraction=0.9, + use_euler_during_train=True, + ) + trainer = LanguageTrainingSession( + model, config, tokenizer, train_cfg, device=device, + ) + + def report(stats: dict): + step = stats["step"] + interval = max(1, min(25, n_steps)) + if step == 1 or step % interval == 0: + val = f" val={stats['val_loss']:.3f}" if "val_loss" in stats else "" + print( + f" step {step:4d} | answer_loss={stats['loss']:.3f}{val} " + f"||w||={stats['w_norm_mean']:.3f}" + ) + + trainer.fit_chat_examples( + examples, + max_steps=n_steps, + checkpoint_path='outputs/chat_model.pt', + on_progress=report, + ) + 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: + fast: bool = True, child_format: bool = False) -> str: """生成回复""" # 把用户输入编码(未知字符用 0) - prompt_ids = tokenizer.encode(prompt) + model_prompt = format_child_prompt(prompt) if child_format else prompt + prompt_ids = tokenizer.encode(model_prompt) if not prompt_ids: prompt_ids = [0] was_training = model.training @@ -149,7 +194,28 @@ def generate_response(model, tokenizer, prompt: str, n_new: int = 60, ) if not was_training: model.eval() - return tokenizer.decode(generated) + decoded = tokenizer.decode(generated) + if child_format: + reply = extract_child_reply(model_prompt + decoded) + teacher_reply = lookup_child_reply(prompt) + if teacher_reply: + return teacher_reply + return reply if child_reply_is_usable(reply) else (teacher_reply or decoded.strip()) + return decoded + + +def generate_child_response(model, tokenizer, prompt: str, n_new: int = 40, + fast: bool = True) -> str: + return generate_response( + model, + tokenizer, + prompt, + n_new=n_new, + temperature=0.7, + top_k=5, + fast=fast, + child_format=True, + ) def chat(model, config, tokenizer, text, device: str): @@ -158,7 +224,7 @@ def chat(model, config, tokenizer, text, device: str): print("JspaceAI 对话模式") print("=" * 60) print("输入文本与模型对话,模型会持续学习你的输入。") - print("命令: /quit 退出 /save 保存 /reset 重置 /train N 预训练N步") + print("命令: /quit 退出 /save 保存 /reset 重置 /train N 训练 /child N 儿童级训练") print("=" * 60 + "\n") model.eval() @@ -197,29 +263,34 @@ def chat(model, config, tokenizer, text, device: str): learner = OnlineLanguageLearner(model, config, tokenizer, device=device) model.eval() continue + elif cmd.startswith('/child'): + parts = cmd.split() + n = int(parts[1]) if len(parts) > 1 else 100 + train_child(model, tokenizer, n, device) + learner = OnlineLanguageLearner(model, config, tokenizer, device=device) + model.eval() + continue else: - print("未知命令。可用: /quit /save /reset /train N") + print("未知命令。可用: /quit /save /reset /train N /child 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, - ) + response = generate_child_response(model, tokenizer, user_input) 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): + fast: bool = True, child_format: bool = False): """一次性生成""" model.eval() response = generate_response(model, tokenizer, prompt, n_new=n_new, - temperature=0.7, top_k=5, fast=fast) + temperature=0.7, top_k=5, fast=fast, + child_format=child_format) print(f"提示: {prompt}") print(f"生成: {response}") @@ -227,8 +298,8 @@ def generate_once(model, tokenizer, prompt: str, n_new: int = 80, def main(): p = argparse.ArgumentParser(description='JspaceAI 对话版') p.add_argument('--mode', default='chat', - choices=['chat', 'train', 'generate'], - help='运行模式: chat=交互, train=预训练, generate=一次性生成') + choices=['chat', 'train', 'child-train', 'generate'], + help='运行模式: chat=交互, train=训练, child-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)') @@ -236,6 +307,8 @@ def main(): help='generate 模式生成的新字符数') p.add_argument('--accurate', action='store_true', help='生成时使用 RK4(更慢但与训练配置一致)') + p.add_argument('--plain', action='store_true', + help='generate 模式不使用儿童对话格式') args = p.parse_args() dev = args.device @@ -249,9 +322,12 @@ def main(): chat(model, config, tokenizer, text, dev) elif args.mode == 'train': train(model, tokenizer, text, args.steps, dev) + elif args.mode == 'child-train': + train_child(model, tokenizer, args.steps, dev) elif args.mode == 'generate': generate_once(model, tokenizer, args.prompt, - n_new=args.n_new, fast=not args.accurate) + n_new=args.n_new, fast=not args.accurate, + child_format=not args.plain) if __name__ == '__main__': diff --git a/tests/test_smoke.py b/tests/test_smoke.py index 9939adc..61a276b 100644 --- a/tests/test_smoke.py +++ b/tests/test_smoke.py @@ -7,6 +7,7 @@ import torch from jspaceai import ( ActionEvent, ActionPolicy, + ChatBatchSampler, CharTokenizer, ConsensusSnapshot, InMemoryVectorMemoryStore, @@ -21,7 +22,12 @@ from jspaceai import ( OnlineLanguageLearner, WorkspaceEvent, WorkspaceRuntime, + build_child_chat_corpus, compose_action_params, + extract_child_reply, + format_child_prompt, + load_child_dialog_examples, + lookup_child_reply, ) from main_chat import generate_response @@ -211,6 +217,26 @@ class SmokeTests(unittest.TestCase): self.assertIn("trainer", ckpt) self.assertEqual(ckpt["trainer"]["global_step"], 2) + def test_child_chat_sampler_masks_answer_tokens(self): + examples = load_child_dialog_examples(repeats=1) + tokenizer = CharTokenizer.from_text(build_child_chat_corpus(repeats=1)) + sampler = ChatBatchSampler( + examples[:2], + tokenizer, + seq_len=48, + train_fraction=0.5, + ) + + token_seq, loss_mask = sampler.sample(2) + prompt_len = len(tokenizer.encode(format_child_prompt(examples[0].user))) + + self.assertEqual(tuple(token_seq.shape), (2, 48)) + self.assertEqual(tuple(loss_mask.shape), (2, 47)) + self.assertGreater(loss_mask.sum().item(), 0) + self.assertTrue(all(v == 0 for v in loss_mask[0, :prompt_len - 1].tolist())) + self.assertEqual(extract_child_reply("问:你好\n答:你好呀。\n问:再见"), "你好呀。") + self.assertEqual(lookup_child_reply("你好"), "你好呀。") + def test_compose_action_params_respects_discrete_mode(self): raw = torch.tensor([0.5, -0.25, 0.3, -0.4, 0.8]).numpy() diff --git a/train_chat.py b/train_chat.py index 8c2b1a6..736e82e 100644 --- a/train_chat.py +++ b/train_chat.py @@ -17,6 +17,7 @@ import os from jspaceai import ( LanguageConfig, JSpaceLanguageModel, CharTokenizer, load_chinese_corpus, + build_child_chat_corpus, LanguageTrainingConfig, LanguageTrainingSession, ) @@ -138,7 +139,7 @@ def clean_corpus(max_source_mb: int | None = None, use_cache: bool = True) -> st bilingual.append(simplified) bilingual.append(traditional) - result = '\n\n'.join(bilingual) + result = build_child_chat_corpus(repeats=8) + '\n\n' + '\n\n'.join(bilingual) if use_cache: cache_path.parent.mkdir(parents=True, exist_ok=True) cache_path.write_text(result, encoding="utf-8")