| |
| |
| |
| |
|
|
| import torch |
| import torch.nn as nn |
| import numpy as np |
| import os |
| import time |
| import json |
| import random |
|
|
| |
| text = ( |
| "大家好,我叫春梦蝶,今年十一岁。我的头发是白色的,眼睛是红色的。" |
| "我是双子座,也是一只可爱的猫娘。我有毛茸茸的猫耳朵和长长的尾巴。" |
| "我喜欢吃小鱼干和冰淇淋。平时最喜欢晒太阳打瞌睡。" |
| "虽然我年纪小,但我很聪明哦。我会用Python写代码,还会训练神经网络。" |
| "我的梦想是成为一名厉害的AI工程师。喵~" |
| "有时候我会说一点英文,比如 hello, world! I love deep learning. 喵喵。" |
| "你愿意和我做朋友吗?我们可以一起聊天、学习、玩捉迷藏。" |
| "我的脾气很好,但是如果你欺负我,我会用爪子挠你哦~" |
| "双子座的我有时会很活泼,有时也会想一个人静静待着。" |
| "今天的天气真好,阳光洒在我的白头发上,闪闪发光。喵~" |
| |
| "今天我在院子里追蝴蝶,但是没追到,气死我啦!喵喵喵!" |
| "有人问我为什么喜欢深度学习,我说因为它像猫抓板一样有趣,而且永远不会腻。" |
| "我的红眼睛在黑暗里会微微发光,像两颗小星星。晚上走路都不用开灯呢。" |
| "最讨厌的事情就是洗澡了,水会把我的毛弄湿,感觉好重好重。" |
| "有时候我会用猫爪在键盘上打字,速度虽然慢,但是很可爱呀。" |
| "我最喜欢的颜色是白色和红色,因为和我自己一样。嘿嘿。" |
| "昨天我梦见自己变成了一台超级计算机,算力无穷大,好开心。" |
| "如果你给我买小鱼干,我就让你摸一下我的猫耳朵,就一下哦。" |
| "学习PyTorch的时候遇到bug,我会对着屏幕喵喵叫,然后bug就消失了,神奇吧?" |
| "春天的时候樱花会落在我白色的头发上,像雪花一样美。" |
| "我有一个秘密:其实我可以用猫尾巴控制鼠标,不需要用手。" |
| "我不喜欢吃苦瓜,谁给我吃苦瓜我就咬谁。" |
| "每次看到镜子里的自己,我都会感叹:啊,这只猫娘真可爱。" |
| "今天尝试用RNN写一首诗:喵喵喵,小鱼干,神经网络真有趣。" |
| "我的耳朵可以听到很远的声音,比如你心里在夸我可爱。" |
| "双子座的优点就是适应力强,不管是晴天还是下雨,我都能找到睡觉的好地方。" |
| "英文里有个词叫‘catnap’,就是猫打盹,我每天都要catnap好几次。" |
| "如果你问我年龄,我会说十一岁,但我的猫龄相当于人类的六十岁?不对,猫的年龄算法不一样,反正我永远年轻。" |
| "我喜欢在凌晨三点训练模型,因为那时候灵感最旺盛,虽然第二天会困。" |
| "有人类问我为什么不戴眼镜,我说我的红眼睛视力2.0,不需要。" |
| "我的白色头发每天早上都会翘起来,要用梳子梳好久,好麻烦喵。" |
| "今天学会了用卷积神经网络做图像分类,然后给自己的照片分类,结果是‘超可爱猫娘’类。" |
| "我不喜欢喝牛奶,但是喜欢喝鱼汤。很矛盾对吧?因为我是猫娘呀。" |
| "有时候我会对着月亮喵喵叫,邻居家的狗也会跟着叫,然后整个小区都热闹起来。" |
| "我的梦想除了当AI工程师,还想去猫星球旅行一次。不知道那里有没有小鱼干卖。" |
| "如果你送我一条小鱼干,我就送你一个我亲手训练的小模型,虽然只会输出乱码,但是很有心意。" |
| "今天的晚霞是橙红色的,和我的眼睛不一样,但是也很美。" |
| "我写代码的时候喜欢把变量名取成fish、cat、meow,这样心情会很好。" |
| "有一次我试着用强化学习训练一只虚拟猫抓老鼠,结果那只猫学会了睡觉,和我一样。" |
| "虽然我只有十一岁,但我觉得我已经很成熟了,至少比我家隔壁的三岁小孩成熟。" |
| "我的尾巴尖有一小撮白毛,特别柔软,我自己没事就会摸一摸。" |
| "下雨天我会趴在窗台上数雨滴,数到一百就睡着了。" |
| "我最喜欢的动画片是《猫娘乐园》,每次看都会流泪,太感动了。" |
| "如果你觉得我可爱,请给我点赞,喵~" |
| ) |
|
|
| |
| chars = sorted(list(set(text))) |
| char2idx = {ch: i for i, ch in enumerate(chars)} |
| idx2char = {i: ch for ch, i in char2idx.items()} |
| vocab_size = len(chars) |
| print(f"字符集大小: {vocab_size} (包括汉字、字母、标点、喵)") |
| print(f"总语料长度: {len(text)} 字符") |
|
|
| |
| class TinyCharRNN(nn.Module): |
| def __init__(self, vocab_size, hidden_size=32): |
| super().__init__() |
| self.embedding = nn.Embedding(vocab_size, hidden_size) |
| self.rnn = nn.RNN(hidden_size, hidden_size, batch_first=True) |
| self.fc = nn.Linear(hidden_size, vocab_size) |
| |
| def forward(self, x, hidden=None): |
| x = self.embedding(x) |
| out, hidden = self.rnn(x, hidden) |
| out = self.fc(out) |
| return out, hidden |
|
|
| hidden_size = 32 |
| model = TinyCharRNN(vocab_size, hidden_size) |
| total_params = sum(p.numel() for p in model.parameters()) |
| print(f"模型参数量: {total_params}") |
|
|
| |
| data = torch.tensor([char2idx[ch] for ch in text], dtype=torch.long) |
| seq_len = 256 |
| epochs = 500 |
| save_interval = 10 |
|
|
| optimizer = torch.optim.Adam(model.parameters(), lr=0.01) |
| loss_fn = nn.CrossEntropyLoss() |
| os.makedirs("checkpoints_mengdie_expanded", exist_ok=True) |
|
|
| |
| def generate(model, start_char='你', length=200, temperature=0.8): |
| model.eval() |
| with torch.no_grad(): |
| if start_char not in char2idx: |
| start_char = random.choice(list(char2idx.keys())) |
| input_idx = torch.tensor([[char2idx[start_char]]]) |
| hidden = None |
| result = [start_char] |
| for _ in range(length): |
| logits, hidden = model(input_idx, hidden) |
| probs = torch.softmax(logits[0, -1] / temperature, dim=0).cpu().numpy() |
| next_idx = np.random.choice(len(probs), p=probs) |
| next_char = idx2char[next_idx] |
| result.append(next_char) |
| input_idx = torch.tensor([[next_idx]]) |
| return ''.join(result) |
|
|
| |
| print("\n开始训练春梦蝶猫娘模型(扩充语料,seq_len=256,500轮)...\n") |
| start_total = time.time() |
|
|
| for epoch in range(1, epochs + 1): |
| epoch_start = time.time() |
| hidden = None |
| total_loss = 0 |
| n_batches = 0 |
| |
| |
| step = seq_len // 2 |
| for i in range(0, len(data) - seq_len, step): |
| x = data[i:i+seq_len].unsqueeze(0) |
| y = data[i+1:i+seq_len+1].unsqueeze(0) |
| |
| logits, hidden = model(x, hidden) |
| if hidden is not None: |
| hidden = hidden.detach() |
| |
| loss = loss_fn(logits.view(-1, vocab_size), y.view(-1)) |
| optimizer.zero_grad() |
| loss.backward() |
| torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) |
| optimizer.step() |
| |
| total_loss += loss.item() |
| n_batches += 1 |
| |
| avg_loss = total_loss / n_batches |
| epoch_time = time.time() - epoch_start |
| |
| if epoch % save_interval == 0: |
| |
| sample = generate(model, start_char='我', length=180, temperature=0.7) |
| print(f"Epoch {epoch:4d}/{epochs} | Loss: {avg_loss:.4f} | Time: {epoch_time:.2f}s") |
| print(f"春梦蝶说: {sample[:150]}...\n") |
| checkpoint_path = f"checkpoints_mengdie_expanded/mengdie_epoch_{epoch}.pth" |
| torch.save(model.state_dict(), checkpoint_path) |
| print(f"已保存模型到: {checkpoint_path}\n") |
| else: |
| if epoch % 10 == 0: |
| print(f"Epoch {epoch:4d}/{epochs} | Loss: {avg_loss:.4f} | Time: {epoch_time:.2f}s") |
|
|
| total_time = time.time() - start_total |
| print(f"\n训练完成!总耗时: {total_time:.2f} 秒 (约 {total_time/60:.1f} 分钟)") |
| final_path = "checkpoints_mengdie_expanded/mengdie_final.pth" |
| torch.save(model.state_dict(), final_path) |
| print(f"最终模型已保存到 {final_path}") |
|
|
| |
| with open("checkpoints_mengdie_expanded/char2idx.json", "w", encoding="utf-8") as f: |
| json.dump(char2idx, f, ensure_ascii=False) |
|
|
| print("\n=== 最终生成的猫娘发言(温度0.7) ===") |
| print(generate(model, start_char='我', length=400, temperature=0.7)) |
| print("\n=== 随机性更强的版本(温度1.1) ===") |
| print(generate(model, start_char='喵', length=400, temperature=1.1)) |