前面学习大模型时,我一直有几个疑问:
- token 到底怎么进入模型?
- 4 层、256 维、4 个注意力头,分别体现在哪里?
- 训练时,究竟是哪行代码改了权重?
- 训练完保存的东西,又怎么拿来回答问题?
光看概念容易绕,所以这次直接做了一个小实验:不下载现成模型,从随机权重开始,用几百条中文短句,训练一个约 321 万参数的 Transformer。
结果有点意思:它从乱续写,变成了能输出像样的短句,但仍然会把"紫色"回答成"黄色"。
这篇不讲"几分钟造出一个 ChatGPT",而是把一个真实的训练过程拆开看。正文讲核心逻辑,文末附完整训练脚本。
1. 先说清楚:这次做的不是微调
这次没有加载 Qwen、DeepSeek 或其他预训练权重,也没有调用云端模型 API。
我们自己指定模型结构,让框架初始化参数,然后用训练资料调整它们。矩阵权重随机初始化,部分偏置和归一化参数按框架默认方式初始化;没有继承已有模型的语言能力。
因此,它是一个从零训练的字符级语言模型,不是 LoRA,也不是蒸馏。
"从零"说的是不使用预训练权重,不是手写 GPU 驱动、矩阵乘法和自动求导。底层计算仍然交给 MLX。
本文使用 MLX 的神经网络组件搭建模型,并直接训练;不需要调用 MLX-LM 的微调入口。可用组件见 MLX 神经网络文档。
2. 这次模型有多小?
本次实测环境是 Apple M3 Max、96 GB 统一内存、macOS 26.5.2,使用原生 arm64 Python 3.12.1 和 MLX 0.32.2。实际运算设备是 GPU。
模型配置如下:
| 配置 | 本次值 | 怎么理解 |
|---|---|---|
| Transformer 层数 | 4 | 数据依次经过 4 层计算 |
| 模型宽度 | 256 | 每个 token 在层间用 256 个数字表示 |
| 每层注意力头数 | 4 | 同层内有 4 个并行的注意力头 |
| 每头宽度 | 64 | 本模型中,256 ÷ 4 = 64 |
| FFN 中间宽度 | 1024 | 每层做 256 → 1024 → 256 的加工 |
| 最大上下文 | 128 | 最多处理 128 个字符 token 的输入窗口 |
| 词表大小 | 112 | 111 个字符,加 1 个 PAD 补齐标记 |
| 参数总量 | 3,212,800 | 约 3.21M,也就是 0.00321B |
这里最容易混淆的是"层"和"头"。左侧是整体结构,右侧展开其中一层:

图中蓝色是注意力相关计算,绿色是 FFN,带"+"的圆圈表示残差相加。层与层前后串联,每层里的头并行计算,所以不是"4 × 4 = 16 层"。注意力头内部包含 Q/K/V 投影和因果注意力计算。
每个头通过各自的投影权重处理输入,得到自己的 64 维表示。这里不是把原始向量机械切成四份,再规定某个头负责"颜色"、另一个头负责"地点";头学到什么,并没有这样的人为分工。
FFN 则对每个位置的表示分别加工。注意力负责在 token 之间交换信息,FFN 本身不直接进行跨 token 的注意力计算。
3. 训练资料从哪里来?
为了先看懂流程,我没有下载大型语料库,而是在代码中组合人名、食物、地点和颜色,生成了 640 条短句。
例如:
text
故事:小白在餐厅。
问:小白在哪里?
答:餐厅。
故事:花朵是紫色的。
问:花朵是什么颜色?
答:紫色。
核心生成方式并不复杂:
python
samples = [
f"故事:{name}在{place}。\n问:{name}在哪里?\n答:{place}。\n"
for name in names
for place in places
]
这批资料分成两部分:576 条用于训练,64 条用于验证。先按完整样本去重,再固定随机种子打乱、划分,两个集合不包含完全相同的样本。
但必须说明:两组资料仍然共用同一套句式。 这种验证只能观察相似句式上的表现,不能拿来证明模型已经理解通用中文。
从零训练一个教学模型,少量合成资料就能跑通流程;从零训练一个真正通用的模型,是另一种规模的数据与计算问题。不能混为一谈。
4. 文字怎样变成 token?
这次采用最容易观察的字符级方案:一个字符对应一个编号,标点和换行也算字符。它不是商业大模型普遍采用的完整分词方案。
词表只从训练集生成:
python
# 0 留给补齐标记,其他编号对应实际字符。
tokens = ["<PAD>"] + sorted(set("".join(train)))
stoi = {char: index for index, char in enumerate(tokens) if index > 0}
本次保存的词表里,餐 对应编号 103。这个编号只是索引,不表示"餐"的权重是 103,也不表示编号大的字更重要。
训练时,把同一句话错开一个位置:
text
原文:小 白 在 餐 厅 。
输入:小 白 在 餐 厅
目标:白 在 餐 厅 。
意思是:看到"小",预测"白";看到"小白",预测"在";看到"小白在",预测"餐"......
实际代码会把不足 128 个位置的部分补齐,并让这些位置不参与 loss。
另外,本实验对故事、问题、答案中的所有有效"下一个字符"位置都计算损失,不是只在答案区域计算。这一点也会影响我们后面对 loss 的理解。
5. 核心模型,其实主要是在组合组件
下面是结构的简化摘录,不是独立运行的完整脚本:
python
# token 编号 → 256 维向量。
self.embedding = nn.Embedding(112, 256)
# 告诉模型字符的先后位置。
self.position = nn.SinusoidalPositionalEncoding(256)
# 4 层,每层 4 个头,FFN 中间宽度 1024。
self.transformer = nn.TransformerEncoder(
num_layers=4,
dims=256,
num_heads=4,
mlp_dims=1024,
dropout=0.0,
norm_first=True,
)
# 对词表中的 112 个 token 分别打分。
self.output = nn.Linear(256, 112, bias=False)
每层里的注意力、FFN、残差连接和归一化,框架已经提供。我们配置层数、宽度等结构参数,再把词嵌入和输出层接起来。
数据真正经过模型时,核心只有几行:
python
hidden = self.embedding(tokens) + self.position(mx.arange(length))
mask = nn.MultiHeadAttention.create_additive_causal_mask(length)
hidden = self.transformer(hidden, mask)
scores = self.output(hidden)
其中,mask 很关键:它禁止当前位置看到后面的字符,防止模型偷看正确答案。
虽然这里的框架组件叫 TransformerEncoder,但我们显式加了因果掩码,把它用于自回归语言建模。这里没有再接一个带交叉注意力的解码器。
输出层得到的是各 token 的原始分数,也叫 logits;它们不是另一个模型给出的人工评分。
6. 训练到底怎样调整权重?
把训练与推理放在一起看,最重要的区别是:训练根据误差更新权重,推理使用已保存的权重继续预测。

左侧循环一次表示一个训练步;右侧循环一次通常选出一个新 token。两者都会经过模型的 4 层,但这两个循环的含义不同。
先定义误差:
python
def loss_fn(model, x, y, mask):
# 比较预测分数与正确 token,得到每个位置的交叉熵。
errors = nn.losses.cross_entropy(model(x), y)
# 只对真实字符位置取平均,不计算补齐位置。
return (errors * mask).sum() / mask.sum()
这里不是拿"预测编号 103"和"正确编号 31"做减法。交叉熵关注的是:模型给正确 token 分配的预测概率有多大。
再进行参数更新。下面省略了运行时限检查和日志:
python
optimizer = optim.AdamW(learning_rate=0.001, weight_decay=0.01)
loss_and_grad = nn.value_and_grad(model, loss_fn)
for step in range(400):
# 每步抽取 16 条样本,不同训练步可能再次抽到同一条。
x, y, mask = make_batch(rng.sample(train, 16), stoi, 128)
# 前向计算与反向传播,得到 loss 和梯度。
loss, grads = loss_and_grad(model, x, y, mask)
# 限制过大的梯度。
grads, _ = optim.clip_grad_norm(grads, 1.0)
# 根据梯度调整权重,再触发实际计算。
optimizer.update(model, grads)
mx.eval(model.parameters(), optimizer.state)
optimizer.update(model, grads) 就是参数更新的关键位置。
梯度可以粗略理解为:参数往哪个方向改变,有助于减小当前误差。AdamW 根据梯度和自己维护的状态更新参数,并不是人工逐个修改几百万个数字。
这里的 400 步是更新权重 400 次,不是 400 层,也不是只处理或生成 400 个 token。
训练时,一批样本中多个位置的预测可以一起计算;因果掩码负责隔离未来信息。它不需要像用户聊天那样,逐字生成整段答案后才开始算误差。
7. 训练后,它真的变好了吗?
先看同一个完整验证集上的 loss:
| 已训练步数 | 验证集 loss |
|---|---|
| 0 | 4.7208 |
| 50 | 0.9483 |
| 100 | 0.6604 |
| 200 | 0.5359 |
| 250 | 0.5400 |
| 400 | 0.4935 |
loss 总体下降,但不保证每一步、每个检查点都下降。
本次完整训练集 loss 从 4.7344 降到 0.4677。数据准备、400 步训练、验证、保存和重载检查这段计时合计约 7.319 秒,不包含环境安装、代码编写、单测和事后复核。
这个时间来自上述硬件上的小规模实验,不是通用性能结论,更不代表"训练一个会聊天的模型只要 7 秒"。
再看实际回答:
| 题目中的事实 | 正确答案 | 模型续写 | 结果 |
|---|---|---|---|
| 小白在餐厅 | 餐厅。 | 餐厅。 | 对 |
| 花朵是紫色的 | 紫色。 | 黄色。 | 错 |
| 椅子是红色的 | 红色。 | 白色。 | 错 |
随后对 64 条验证题逐条生成答案,去除首尾空白后做精确匹配,答对 26 条,即 40.625%。这是额外的生成评测,不是上表的 loss。
为什么 loss 已经降得很多,答案还是会错?
因为本实验的 loss 覆盖整个样本,模型学会常见句式、标点、固定措辞,也能降低平均误差,但这不等于它一定能正确提取每个问题需要的信息。
这组结果说明:它学到了一部分短句规律,但还没有可靠掌握这类问答。 目前证据不能进一步断言"只是再训练多少步就能全部答对"。
验证集在训练过程中已被查看,而且与训练集共用模板,所以这里也不把 26/64 冒称为独立测试集成绩或通用能力分数。
8. 保存后,怎样继续生成文字?
本次正式运行保存了这些文件:
| 文件 | 内容 |
|---|---|
weights.safetensors |
训练后的参数权重 |
config.json |
层数、宽度、头数、上下文等配置 |
tokenizer.json |
字符和编号的对应关系 |
parameter_shapes.json |
每块参数的名称与形状 |
train.txt、validation.txt |
这次使用的两组自制资料 |
metrics.json |
损失、耗时、生成前后对比、重载检查等 |
仅权重文件约 12.86 MB;整个正式运行目录约 12.9 MB。代码实测的 MLX 内存峰值约 513 MB,不能把它理解成整个系统或 Python 进程的总内存,也不能把权重文件大小当成运行内存。
权重、词表、配置需要配套加载:同一个编号在不同词表里可能对应不同字符,不能随意混用。
重载时,我们创建全新模型实例,再从文件读取权重。本次固定输入的 logits 最大绝对差为 0,验证 loss 和生成文本也一致。
输入这句话时:
text
故事:小白在餐厅。
问:小白在哪里?
答:
它一共有 21 个字符 token。实测流程是:
text
21 个 token 编号
→ 21 × 256 的向量表示
→ 依次经过 4 层
→ 取最后一个位置的输出分数
→ 选中 token 103,也就是"餐"
→ 把"餐"追加进去,再预测后面的字
最终续写为 餐厅。,接着生成换行并停止。
本实验选择最高分 token,叫贪心生成,不是每次随机抽一个字。测试过程中没有执行优化器更新,权重不会因为这次提问而改变。
为了让实现容易阅读,这个版本没有 KV Cache,每生成一个字符都会重新计算窗口内的上下文。因此,它适合教学,不是高性能推理服务。
9. 自己动手:先训练,再输入问题
下面给出从新目录开始的步骤。命令面向 macOS 的 zsh/bash;其他硬件或系统版本请先检查 MLX 官方安装要求。
先创建一个全新的实验目录,已有同名目录时不要直接覆盖:
bash
mkdir mlx-from-scratch-demo
cd mlx-from-scratch-demo
python3 -m venv .venv
source .venv/bin/activate
python -m pip install "mlx==0.32.2"
把文末完整脚本保存为当前目录的 train.py,先做两步冒烟,再正式训练:
bash
python train.py --steps 2 --output runs/smoke-2
python train.py --steps 400 --output runs/first-400
同名输出目录已存在时,脚本会报错,不会覆盖旧结果。再次训练时换一个目录名即可。
脚本每步检查训练耗时,达到 240 秒时停止继续更新,随后进行验证和保存。请查看输出里的实际 steps 或 metrics.json 的 steps_completed,不要把因超时提前结束当成完成了 400 步。本文运行时还在外层设置了 300 秒进程超时,这与训练循环的 240 秒检查不是同一层限制。
训练结束后,加载自己的权重测试:
bash
python train.py --load runs/first-400 \
--prompt $'故事:小白在餐厅。\n问:小白在哪里?\n答:'
本次权重的输出是:
text
故事:小白在餐厅。
问:小白在哪里?
答:餐厅。
也可以运行下面这段交互循环,先输入故事,再输入问题。它复用已保存的模型,不会再次训练:
bash
python -c '
from train import load_checkpoint, generate
model, tokens = load_checkpoint("runs/first-400")
print("每轮输入故事和问题;输入 /exit 退出。")
while True:
try:
story = input("故事 > ").strip()
if story == "/exit":
break
if not story:
continue
question = input("问题 > ").strip()
if question == "/exit":
break
if not question:
continue
prompt = f"故事:{story}\n问:{question}\n答:"
print("模型 > " + generate(model, tokens, prompt).strip())
except ValueError as error:
print(error)
except (EOFError, KeyboardInterrupt):
break
'
先试这组输入:故事填 小白在餐厅。,问题填 小白在哪里?。
需要注意:
- 它是句式续写模型,不是对话助手;直接输入"你好"并不能保证得到问候。
- 它只认识这次词表里的字符,遇到陌生字符会报错。
- 最多生成 32 个新字符,遇到换行也会结束;生成窗口超过 128 个字符时,只保留最近的部分。
- 即使提示看起来很简单,也可能答错;版本、硬件和数值计算差异也可能影响复现实验的具体结果。
10. 这次真正学到了什么?
我觉得最有价值的不是"得到了一个能用的聊天模型",而是把几个原来抽象的概念连起来了:
- 模型结构由配置决定,训练主要改变参数,而不是自动决定变成多少层。
- token 是编号,进入模型后变成向量,再通过多层计算得到下一 token 的分数。
- 训练是"预测 → 算误差 → 求梯度 → 更新权重",不是人工给每个参数评分。
- 推理仍然执行模型计算,但不执行优化器更新。
- loss 下降、能输出句子、答题正确,是不同层面的证据,不能混在一起。
一个 321 万参数的小实验,已经足够把这些步骤真正跑一遍。至于更好的问答能力,需要继续设计数据、训练与评测实验,而不是先把这次有限的结果包装成成功。
附录:完整训练脚本
保存为 train.py。以下是本次实验实际使用的完整脚本;保留了输入校验、资源边界、日志、模型保存和重载检查。上文的 26/64 来自训练结束后的额外答案精确匹配评测,不由这段脚本自动汇总。
python
"""从随机权重训练中文字符级小模型;只使用自制句式,不访问网络。"""
import argparse
import hashlib
import json
import math
from pathlib import Path
import random
import time
import mlx.core as mx
import mlx.nn as nn
import mlx.optimizers as optim
from mlx.utils import tree_flatten
def make_corpus():
# 用自制的人名和事物组合教学样本,不读取任何私人文件。
names = "小明 小红 小华 小林 小雨 小雪 小东 小西 小南 小北 小安 小乐 小青 小白 小夏 小秋".split()
foods = "苹果 香蕉 西瓜 葡萄 草莓 橙子 桃子 梨子 面包 米饭 面条 饺子 牛奶 豆浆 土豆 玉米".split()
places = "公园 学校 图书馆 厨房 客厅 花园 书店 商店 教室 操场 山上 河边 车站 广场 餐厅 家里".split()
objects = "书包 杯子 帽子 衣服 鞋子 铅笔 书本 椅子 桌子 汽车 雨伞 气球 花朵 盒子 水杯 皮球".split()
colors = "红色 蓝色 黄色 绿色 白色 黑色 紫色 橙色".split()
# 每条样本同时包含事实、问题和答案,让模型学习简单的续写与信息提取。
samples = [f"故事:{name}喜欢{food}。\n问:{name}喜欢什么?\n答:{food}。\n" for name in names for food in foods]
samples += [f"故事:{name}在{place}。\n问:{name}在哪里?\n答:{place}。\n" for name in names for place in places]
samples += [f"故事:{obj}是{color}的。\n问:{obj}是什么颜色?\n答:{color}。\n" for obj in objects for color in colors]
# 先按完整样本去重、固定打乱,再留出十分之一作验证集。
samples = sorted(set(samples))
random.Random(20260907).shuffle(samples)
split = len(samples) // 10
return samples[split:], samples[:split]
def make_batch(samples, stoi, context):
# 每个位置用当前字符预测下一个字符;0 专用于补齐,不代表实际汉字。
inputs, targets, masks = [], [], []
for sample in samples:
# 超长和陌生字符显式报错,避免静默改写训练资料。
if not 2 <= len(sample) <= context + 1:
raise ValueError("样本必须有至少两个字符,且预测长度不能超过上下文")
if not set(sample) <= stoi.keys():
raise ValueError("出现训练词表以外的字符")
ids = [stoi[char] for char in sample]
size = len(ids) - 1
# 输入与标签错开一位;补齐位置不能贡献 loss。
inputs.append(ids[:-1] + [0] * (context - size))
targets.append(ids[1:] + [0] * (context - size))
masks.append([True] * size + [False] * (context - size))
return mx.array(inputs), mx.array(targets), mx.array(masks)
class TinyLM(nn.Module):
def __init__(self, config):
super().__init__()
# 配置定义模型的形状;具体权重由 MLX 随机初始化,再由训练调整。
self.context = config["context"]
self.embedding = nn.Embedding(config["vocab_size"], config["dims"])
# 位置编码告诉模型字的先后次序,本实现没有可训练的位置参数。
self.position = nn.SinusoidalPositionalEncoding(config["dims"])
# 框架在每层封装多头注意力、FFN、残差与归一化。
self.transformer = nn.TransformerEncoder(
config["layers"], config["dims"], config["heads"],
mlp_dims=config["ffn_dims"], dropout=0.0, norm_first=True,
)
# LM Head 为词表中的每个 token 计算一个分数。
self.output = nn.Linear(config["dims"], config["vocab_size"], bias=False)
def __call__(self, tokens):
# 拒绝空序列及超出设计上下文的输入。
length = tokens.shape[1]
if not 1 <= length <= self.context:
raise ValueError("输入长度必须在上下文范围内")
# 将数字编号变成向量,并叠加对应位置的信息。
hidden = self.embedding(tokens) + self.position(mx.arange(length))
# 尽管框架组件叫 Encoder,这个因果掩码使它只能看当前及此前字符。
mask = nn.MultiHeadAttention.create_additive_causal_mask(length)
hidden = self.transformer(hidden, mask)
return self.output(hidden)
def loss_fn(model, x, y, mask):
# 交叉熵衡量下一个字符的预测误差,仅平均真实字符的位置。
per_token = nn.losses.cross_entropy(model(x), y)
return (per_token * mask).sum() / mask.sum()
def evaluate(model, samples, stoi, context, batch_size=16):
# 验证只前向计算,不求梯度、不修改权重;按真实 token 数加权。
total_loss, total_tokens = 0.0, 0
model.eval()
for start in range(0, len(samples), batch_size):
x, y, mask = make_batch(samples[start:start + batch_size], stoi, context)
count = int(mask.sum().item())
total_loss += float(loss_fn(model, x, y, mask).item()) * count
total_tokens += count
return total_loss / total_tokens
def generate(model, tokens, prompt, max_new_tokens=32):
# 教学版使用确定性贪心生成;不使用 KV Cache,每次重算窗口内的上下文。
stoi = {char: index for index, char in enumerate(tokens) if index > 0}
if not prompt or not set(prompt) <= stoi.keys():
raise ValueError("提示不能为空,且只能使用这个小词表已有的字符")
ids = [stoi[char] for char in prompt]
result = ""
model.eval()
for _ in range(max_new_tokens):
# 超出窗口时只取最后 context 个字符,不自动扩展模型上下文。
logits = model(mx.array([ids[-model.context:]]))[0, -1]
# PAD 只是训练补齐标记,禁止生成;取最高分而不使用随机采样。
next_id = 1 + int(mx.argmax(logits[1:]).item())
char = tokens[next_id]
result += char
ids.append(next_id)
# 数据中答案以换行结尾;换行或生成上限都会结束输出。
if char == "\n":
break
return result
def load_checkpoint(directory):
# 权重必须与原词表编号和形状配置一起加载。
directory = Path(directory)
config = json.loads((directory / "config.json").read_text(encoding="utf-8"))
tokens = json.loads((directory / "tokenizer.json").read_text(encoding="utf-8"))["tokens"]
model = TinyLM(config)
model.load_weights(str(directory / "weights.safetensors"))
mx.eval(model.parameters())
model.eval()
return model, tokens
def train_run(args):
# 所有日志、数据与权重仅写入全新的输出目录,拒绝覆盖既有实验。
output = args.output
output.mkdir(parents=True, exist_ok=False)
started = time.monotonic()
mx.random.seed(7)
rng = random.Random(7)
train, validation = make_corpus()
# 运行时也检查样本隔离,避免未来改语料时引入完整样本泄漏。
if set(train) & set(validation):
raise ValueError("训练集与验证集不能包含相同完整样本")
# 词表只由训练集创建;验证集不用于拟合词表或更新权重。
tokens = ["<PAD>"] + sorted(set("".join(train)))
stoi = {char: index for index, char in enumerate(tokens) if index > 0}
config = dict(vocab_size=len(tokens), layers=4, dims=256, heads=4, ffn_dims=1024, context=128)
model = TinyLM(config)
mx.eval(model.parameters())
# 实际统计所有可训练数组,不用层数或文件大小猜参数量。
parameters = tree_flatten(model.trainable_parameters())
parameter_count = sum(value.size for _, value in parameters)
parameter_shapes = {name: list(value.shape) for name, value in parameters}
initial_embedding = mx.array(model.embedding.weight)
mx.eval(initial_embedding)
# 固定选验证集中的三条,只展示题干到"答:",不把答案喂给生成函数。
examples = validation[:3]
prompts = [example.rsplit("答:", 1)[0] + "答:" for example in examples]
expected = [example.rsplit("答:", 1)[1].strip() for example in examples]
before_samples = [generate(model, tokens, prompt) for prompt in prompts]
before_train_loss = evaluate(model, train, stoi, config["context"])
before_val_loss = evaluate(model, validation, stoi, config["context"])
print(json.dumps({"stage": "before", "parameters": parameter_count, "vocab_size": len(tokens),
"train_samples": len(train), "validation_samples": len(validation),
"val_loss": before_val_loss, "samples": before_samples}, ensure_ascii=False), flush=True)
# AdamW 对全部权重做更新;梯度裁剪限制偶发大步长。
optimizer = optim.AdamW(learning_rate=0.001, weight_decay=0.01)
loss_and_grad = nn.value_and_grad(model, loss_fn)
history = []
steps_completed = 0
stop_reason = "steps"
for step in range(1, args.steps + 1):
# 给最终验证与保存留出时间,总运行时由入口额外设置进程上限。
if time.monotonic() - started >= 240:
stop_reason = "training_time_limit"
break
model.train()
x, y, mask = make_batch(rng.sample(train, 16), stoi, config["context"])
loss, grads = loss_and_grad(model, x, y, mask)
grads, norm = optim.clip_grad_norm(grads, 1.0)
mx.eval(loss, norm)
# 非有限数时立刻停止,避免保存损坏的训练结果。
if not math.isfinite(float(loss.item())) or not math.isfinite(float(norm.item())):
raise RuntimeError("训练出现非有限 loss 或梯度")
optimizer.update(model, grads)
mx.eval(model.parameters(), optimizer.state)
steps_completed = step
# 只记录有意义的采样点,验证使用同一组完整留出样本。
if step == 1 or step % 50 == 0 or step == args.steps:
row = dict(step=step, train_loss=float(loss.item()),
val_loss=evaluate(model, validation, stoi, config["context"]),
elapsed_seconds=round(time.monotonic() - started, 3))
history.append(row)
print(json.dumps(row, ensure_ascii=False), flush=True)
# 训练结束后,再测完整验证集和相同提示的生成结果。
final_val_loss = evaluate(model, validation, stoi, config["context"])
final_train_loss = evaluate(model, train, stoi, config["context"])
after_samples = [generate(model, tokens, prompt) for prompt in prompts]
changed = not bool(mx.array_equal(initial_embedding, model.embedding.weight).item())
# 报告词嵌入这个具体参数张量的变化量,不冒称全模型参数差值。
embedding_delta_l2 = float(mx.sqrt(mx.sum((initial_embedding - model.embedding.weight) ** 2)).item())
model.save_weights(str(output / "weights.safetensors"))
# 保存必要的配置、词表及原始样本,使同一次实验可审计、可重载。
payloads = {"config.json": config,
"tokenizer.json": dict(type="character", pad_id=0, fitted_on="train", tokens=tokens),
"parameter_shapes.json": parameter_shapes}
for name, value in payloads.items():
(output / name).write_text(json.dumps(value, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
(output / "train.txt").write_text("\n".join(train), encoding="utf-8")
(output / "validation.txt").write_text("\n".join(validation), encoding="utf-8")
# 构造全新模型并严格加载;同时核对 logits、验证损失与生成文本。
reloaded, reloaded_tokens = load_checkpoint(output)
probe, _, _ = make_batch(validation[:2], stoi, config["context"])
reload_max_abs_diff = float(mx.max(mx.abs(model(probe) - reloaded(probe))).item())
logits_equal = reload_max_abs_diff <= 1e-6
reload_loss = evaluate(reloaded, validation, stoi, config["context"])
reload_samples = [generate(reloaded, reloaded_tokens, prompt) for prompt in prompts]
samples_equal = after_samples == reload_samples
if not logits_equal or not samples_equal or abs(reload_loss - final_val_loss) > 1e-6:
raise RuntimeError("权重重载验证未通过")
# 记录损失是每个有效字符的自然对数交叉熵,不是准确率。
metrics = dict(
source="synthetic_chinese_templates", initialization="random", device=str(mx.default_device()),
seed=7, split_seed=20260907, parameter_count=parameter_count, vocab_size=len(tokens),
train_samples=len(train), validation_samples=len(validation),
train_tokens=sum(len(sample) - 1 for sample in train),
validation_tokens=sum(len(sample) - 1 for sample in validation),
train_sha256=hashlib.sha256("\n".join(train).encode()).hexdigest(),
validation_sha256=hashlib.sha256("\n".join(validation).encode()).hexdigest(),
steps_requested=args.steps, steps_completed=steps_completed, stop_reason=stop_reason,
learning_rate=0.001, weight_decay=0.01, batch_size=16, gradient_clip_norm=1.0,
before_train_loss=before_train_loss, final_train_loss=final_train_loss,
before_val_loss=before_val_loss, final_val_loss=final_val_loss, reload_val_loss=reload_loss,
embedding_delta_l2=embedding_delta_l2, reload_max_abs_diff=reload_max_abs_diff,
weights_changed=changed, reload_logits_equal=logits_equal, reload_samples_equal=samples_equal,
samples=[dict(prompt=p, expected=e, before=b, after=a) for p, e, b, a in zip(prompts, expected, before_samples, after_samples)],
history=history, peak_mlx_memory_bytes=mx.get_peak_memory(),
elapsed_seconds=round(time.monotonic() - started, 3),
limitations="仅合成句式内的字符预测实验;不等于通用聊天、阅读理解或独立测试集成绩。",
)
# 统计落盘总量,包含 metrics 自身;超出资源边界时显式报错。
metrics_path = output / "metrics.json"
for _ in range(3):
metrics["artifact_bytes"] = sum(path.stat().st_size for path in output.iterdir() if path.is_file())
metrics_path.write_text(json.dumps(metrics, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
if metrics["artifact_bytes"] > 500_000_000:
raise RuntimeError("产物超过 500 MB 限制")
print(json.dumps({"stage": "done", "output": str(output), "val_loss": final_val_loss,
"steps": steps_completed, "seconds": metrics["elapsed_seconds"],
"samples": after_samples, "reload_verified": True}, ensure_ascii=False), flush=True)
def main():
# 仅提供教学需要的训练步数、输出位置和权重加载入口。
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--steps", type=int, default=300)
parser.add_argument("--output", type=Path, default=Path(__file__).parent / "runs" / time.strftime("%Y%m%d-%H%M%S"))
parser.add_argument("--load", type=Path)
parser.add_argument("--prompt", default="故事:小明喜欢苹果。\n问:小明喜欢什么?\n答:")
args = parser.parse_args()
# 上限固定为 500 步,拒绝无意中启动长训练。
if not 1 <= args.steps <= 500:
parser.error("--steps 必须在 1 到 500 之间")
if args.load:
# 推理只读指定 checkpoint,不覆盖任何训练产物。
model, tokens = load_checkpoint(args.load)
print(args.prompt + generate(model, tokens, args.prompt), end="\n")
else:
train_run(args)
if __name__ == "__main__":
main()