读懂数据之后,这个手写数字项目最值得看的部分就是 C 语言实现本身。它没有依赖深度学习框架,而是用一个非常直接的多分类 softmax 模型,把 784 维输入像素映射到 10 个数字类别上。
这类实现很适合训练“把模型公式翻译成代码”的能力。你会看到:权重矩阵怎么定义、softmax 概率怎么计算、交叉熵损失怎么累计、梯度下降怎么一步步更新参数。
一、模型结构其实很简单
项目里最核心的参数只有两组:
- W[10][784]:10 个类别各自对应一组长度为 784 的权重
- b[10]:10 个类别的偏置项
对一条输入样本 x 来说,程序先计算每个类别的线性分数:
z[k] = b[k];
for (int j = 0; j < FEATURES; j++) {
z[k] += W[k][j] * x[j];
}
这一步得到的是 10 个 logits,也就是每个类别当前的原始分数。
二、softmax 把分数变成概率
线性分数本身不方便直接解释成“属于某个数字的概率”,所以程序接着用 softmax 做归一化:
p[i] = exp(z[i] - max_z);
sum += p[i];
...
p[i] /= sum;
这里减去 max_z 是为了数值稳定,避免指数计算时数值太大。softmax 之后,10 个类别的概率会加起来等于 1,程序再选概率最大的类别作为当前预测值。
三、训练循环在做什么
当前项目设置了 20 轮训练,学习率是 0.01。每轮都会遍历训练集中的全部样本:
- 计算 10 个类别的 logits
- 做 softmax,得到概率分布
- 根据真实标签计算误差
- 用误差更新权重和偏置
更新规则写得很直接:
double error = p[k] - (k == y_train[i] ? 1.0 : 0.0);
for (int j = 0; j < FEATURES; j++) {
W[k][j] -= LEARNING_RATE * error * X_train[i][j];
}
b[k] -= LEARNING_RATE * error;
如果你学过逻辑回归或多分类线性模型,会发现这套写法本质上就是 softmax 回归的随机梯度下降版本。它不花哨,但很适合练基本功。
四、输出里最该看哪几个指标
这个项目训练时会打印每轮的损失和训练准确率,训练结束后还会输出训练集准确率和混淆矩阵。对初学者来说,这几项最重要:
- loss:有没有持续下降
- accuracy:分类正确比例有没有逐步提升
- 混淆矩阵:哪些数字最容易互相混淆
如果某几类长期互相错分,通常说明这几个数字的局部形状更接近,或者当前线性模型的表达能力已经接近上限。
五、它是怎么生成 submission.csv 的
训练结束后,程序会逐条读取 test.csv,对每条样本调用一次 predict_one,再写成:
ImageId,Label
1,7
2,2
3,1
...
这就是最终的 submission.csv。从工程角度看,这一步很关键,因为它把“训练代码”真正变成了一个能处理未知输入并导出结果的完整项目。
六、如何在本地运行
站点下载区已经放好了源码、训练集压缩包和测试集压缩包。当前版本直接在源码所在目录读取 train.csv 和 test.csv:
unzip train.csv.zip
unzip test.csv.zip
gcc digit_softmax_classifier.c -lm -O2 -o digit_classifier
./digit_classifier
正常情况下,程序会依次输出:
- 训练样本数和测试样本数
- 每一轮训练的 loss 和 accuracy
- 训练集准确率与混淆矩阵
- 生成
submission.csv的提示
七、这个 C 项目目前的边界在哪里
它已经足够完成一个完整的多分类练习,但也有很清楚的边界:
- 模型仍然是线性的,没有卷积层或更复杂的表示能力
- 当前训练集准确率高,并不等于线上泛化一定最好
- 没有单独划出验证集做调参
- 没有 mini-batch、正则化或更细的学习率调度
这些都不是缺点,而是后续扩展空间。先把一份能跑通、能解释、能导出结果的基础实现做好,本身就很有价值。
八、下一步怎么继续
如果你想先试交互演示,再回来看源码,可以继续打开 实验台里的手写数字标签页。浏览器版不会直接跑完整训练,而是加载一份预训练的轻量 softmax 权重,让你可以手绘数字、看预测概率和样本效果。
源码、压缩数据、样例提交文件和浏览器模型文件都已经放到 下载页。如果你还没看前一篇,建议补读 手写数字数据结构文章,这样这份 C 代码里的每个数组就更容易对上数据来源。
九、手算一条样本的更新方向
理解这份 C 代码时,最关键的是看懂 p[k] - y[k] 这个误差项。假设真实标签是 7,当前模型却给类别 3 更高概率,那么类别 3 的误差为正,类别 7 的误差为负。梯度下降会压低类别 3 在当前像素上的权重,同时抬高类别 7 对这些像素的响应。
真实标签: 7
当前概率: p[3] = 0.62, p[7] = 0.21
更新方向: 降低类别 3 的相关权重,提高类别 7 的相关权重
这就是 softmax 回归可解释的地方:每次更新都和当前样本的像素分布直接相关。某个数字经常被错分成另一个数字,混淆矩阵会把这种模式暴露出来;回到权重矩阵看对应类别,也能理解模型为什么偏向某些形状。
直接照公式写 softmax,一定会溢出
第二节给出的 softmax 公式是数学定义,但照着它直接写 C 代码,遇到稍大的分数就会得到 nan。这是这类项目最经典的一个坑,值得单独说清楚。
问题出在 exp() 上。float 能表示的最大值约 3.4e38,而 exp(89) 就已经超过它了;double 撑到 exp(710) 左右也会溢出成 inf。一旦分子分母都变成 inf,相除得到的就是 nan,而 nan 会顺着反向传播污染所有权重,整个模型一步之内彻底报废。
手写数字这个任务里,输入是 784 维、初期权重又没有约束,分数轻易就能跑到几百,所以这不是罕见边界情况,而是几乎必然发生。
修法基于一个恒等式:softmax 的结果对所有分数同时减去一个常数是不变的。
softmax(z_i) = exp(z_i) / Σ exp(z_j)
= exp(z_i - C) / Σ exp(z_j - C) 对任意常数 C 成立
取 C 等于这一组分数里的最大值,那么所有指数的参数都 ≤ 0,exp() 的结果落在 (0, 1] 区间,永远不会溢出。同时分母里至少有一项等于 1,也就不会出现除以 0。
void softmax(double *z, int n) {
double max = z[0];
for (int i = 1; i < n; i++)
if (z[i] > max) max = z[i]; /* 先找最大值 */
double sum = 0.0;
for (int i = 0; i < n; i++) {
z[i] = exp(z[i] - max); /* 平移之后再取指数 */
sum += z[i];
}
for (int i = 0; i < n; i++)
z[i] /= sum;
}
代价只是多遍历一次十个元素,完全可以忽略。
同一类问题在计算交叉熵损失时会以另一种形式出现:log(p) 在 p 非常接近 0 时会返回 -inf。稳妥的做法是加一个极小的下限:
loss -= log(p[label] < 1e-15 ? 1e-15 : p[label]);
调试时有个很实用的判据:如果 loss 在某一轮突然变成 nan,几乎一定是溢出而不是学习率太大。学习率过大表现为 loss 剧烈震荡或缓慢发散,是有过程的;而溢出是一步到位——上一轮还是正常数值,下一轮直接 nan,之后永远是 nan。看到这个模式就直接去查 exp() 和 log(),不用先怀疑超参数。
十、验证改动没有破坏项目
如果你修改了学习率、训练轮次或数据读取逻辑,不要只看程序是否能编译。至少应该做三类验证:
- 编译验证:使用
gcc -Wall -Wextra检查明显的类型和数组问题。 - 训练验证:确认 loss 大体下降,accuracy 不应长期停在随机水平附近。
- 输出验证:确认
submission.csv行数、表头和标签范围仍然正确。
对于教学项目来说,这些验证比追求一次最高分更重要。只要你能解释每个指标为什么变化,就已经从“运行别人的代码”前进到了“能维护自己的实验”。
十一、Softmax 训练审计表
这份 C 项目的价值在于可解释和可复现。下面的表格把数据读取、数值稳定、训练指标和输出文件放到同一个审计框架里,方便读者判断一次改动到底提升了模型,还是只让代码“看起来能跑”。
| 审计项 | 应该检查什么 | 常见失败模式 | 修复方向 |
|---|---|---|---|
| CSV 读取 | 样本数、字段数、标签范围和像素归一化。 | 表头未跳过、标签列错位或像素未除以 255。 | 在训练前打印前几行解析结果和特征范围。 |
| Softmax 稳定性 | 是否先减去 max_z,概率和是否接近 1。 |
exp() 溢出导致 NaN loss。 |
保留稳定 softmax,并在异常时打印 logits。 |
| 训练趋势 | loss 是否下降,accuracy 是否高于随机水平。 | 长期停在 10% 附近,说明更新或标签可能错了。 | 降低学习率,检查 p[k] - y[k] 更新方向。 |
| 提交文件 | 行数、表头、ImageId 顺序和标签范围。 | 训练正常但导出格式不符合提交要求。 | 把输出验证作为独立步骤,不和训练日志混在一起。 |