用 C 实现手写数字 Softmax 分类器:从 784 维像素到 submission.csv
用 C 实现手写数字 Softmax 分类器:从 784 维像素到 submission.csv
站内搜索
直接问 AI

用 C 实现手写数字 Softmax 分类器:从 784 维像素到 submission.csv

读懂数据之后,这个手写数字项目最值得看的部分就是 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。每轮都会遍历训练集中的全部样本:

  1. 计算 10 个类别的 logits
  2. 做 softmax,得到概率分布
  3. 根据真实标签计算误差
  4. 用误差更新权重和偏置

更新规则写得很直接:

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.csvtest.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 顺序和标签范围。 训练正常但导出格式不符合提交要求。 把输出验证作为独立步骤,不和训练日志混在一起。

发表回复

向下探索