三元字符语言模型两种实现损失不同,二者是否等价?
字符级三元语言模型两种实现方案的对比与疑问解答
背景
我正在跟随Andrej Karpathy的makemore系列实现字符级三元语言模型,目前有两种实现方案,想确认二者在数学上是否等价,还是本质不同的模型。
实现1:直接使用27×27×27权重张量
W = torch.randn((27, 27, 27), requires_grad=True) for k in range(200): logits = W[xs1, xs2] counts = logits.exp() probs = counts / counts.sum(1, keepdim=True) loss = -probs[torch.arange(num), ys].log().mean() W.grad = None loss.backward() W.data += -50 * W.grad # 说明:xs1和xs2是字符索引的整数张量,W[xs1, xs2]直接索引3D权重张量得到形状为(N, 27)的logits
实现2:拼接独热向量搭配54×27权重矩阵
W= torch.randn((54, 27), requires_grad=True) for k in range(200): xenc1 = F.one_hot(xs1, num_classes=27).float() xenc2 = F.one_hot(xs2, num_classes=27).float() xenc = torch.cat([xenc1, xenc2], dim=1) logits = xenc @ W loss = F.cross_entropy(logits, ys) W.grad = None loss.backward() W.data -= 50 * W.grad.data
我的理解
- 实现1拥有27×27×27=19683个参数,每个(字符1,字符2)对都有一套完全独立的27个权重。
- 实现2拥有54×27=1458个参数,由于拼接操作与矩阵乘法,字符1和字符2的贡献是相加的——字符1选取W[0:27]行,字符2选取W[27:54]行,二者求和。
- 因此我认为这并非等价模型:实现1表达能力更强但需要更多数据,实现2做出了加法假设,在数据较少时泛化性更好。
疑问解答
1. 我的理解是否正确,即这两个模型本质不同且数学上不等价?
你的理解完全正确,二者本质不同且数学上不等价:
- 实现1是无结构约束的三元组模型,每个(c1,c2)上下文对对应独立的输出权重,参数空间没有任何限制,能拟合任意三元组概率分布,表达能力极强。
- 实现2是线性加性模型,引入了"两个上下文字符对输出的贡献可加"的强假设,参数空间是实现1的严格子集(无法覆盖实现1的所有可能权重配置),表达能力远弱于实现1。
2. 在Karpathy的makemore1视频所用的小型数据集names.txt(约32k单词)上,哪个模型会收敛到更低的损失?
在训练损失上,实现1会收敛到更低的值:
- 实现1参数更多,拟合能力更强,能捕捉数据中所有细微的三元组统计规律,甚至包括噪声;实现2受限于参数规模和加性假设,无法完全拟合数据中的所有细节。
- 不过要注意,若数据集过小,实现1可能出现过拟合,导致测试集损失高于实现2,但names.txt对应的三元组样本量达几十万级别,对于19k参数的实现1来说数据量足够,训练损失会显著低于实现2。
3. 实现1本质上是否是一个用梯度下降训练的查找表,等价于计数模型?
实现1等价于带平滑的计数模型,但和朴素计数模型有区别:
- 朴素计数模型直接统计三元组出现频率,未出现的三元组概率为0,存在零概率问题;而实现1通过梯度下降训练交叉熵损失,初始化的随机权重给未出现的三元组赋予了非零初始概率,训练过程中会根据数据调整,相当于给计数模型加了平滑。
- 当数据量无限大时,实现1的结果会趋近于朴素计数模型(最大化似然等价于统计频率),但在有限数据下,它是带有平滑效果的版本,和纯计数模型不完全相同。
内容的提问来源于stack exchange,提问作者Tilak Soni
相关产品推荐
相关产品推荐

