手写数字识别器的设计与实现

问题分析

手写数字的识别本质上是一个经典的图像样本分类问题。对于这类问题,统计学上已经有许多经典的解决方法,如支持向量机、kNN聚类、人工神经网络等等。本文选择使用Matlab实现一个简单的基于矩阵运算的神经网络,并将其用于手写数字的识别。通过增大输入层神经元的数量,我们可以不抽取图像的特征参数,而是直接将图像本身输入网络的输入层来进行训练,从而减轻了工作量。为了训练该网络,我们选用著名的MNIST数据集,其中包含60000个训练样本,以及10000个测试样本。

当神经网络训练完成以后,我们可以将其用于识别实际图片中的手写数字。为此,我们需要一种有效的算法,来将图片中手写数字的部分分离出来,经过适当的预处理后转换为神经网络所支持的样式,然后再进行识别。因此,本文提出了一种基于连通分量思想的有效算法,来实现笔迹的分离操作。

该识别器的所有编程实现均基于Matlab,点击此处访问Github上的对应repo,可以获取所有源代码(不包含训练数据与训练结果)。

神经网络

拓扑结构

为了简化问题,同时也是考虑到手写数字训练样本的分辨率并不高(28x28),因此本文忽略了图像特征向量的提取,而是直接将图像本身作为向量输入网络。网络的拓扑结构大体如下图所示:

神经网络拓扑结构

其中,输入层包含784个单元,即为输入数字图像的像素个数;中间层包含50个单元,这是根据经验选定的值,应该足以记录手写数字所包含的特征信息;而输出层包含10个单元,分别对应于0~9这10个数字,用来输出某个图像对应于某个数字的概率。

数据与训练

本文使用反向传播算法(Error Back Propagation)来训练神经网络,使用的训练集是MNIST数据集的60000个手写数字样本,测试集是10000个手写数字样本。算法的大体流程是:

  1. 将网络输入层与隐含层、隐含层与输出层之间的权重进行随机初始化;
  2. 构造代价函数,用于衡量对于所有训练数据,当前网络的分类效果与预期效果的差异,并计算反向传播的梯度;
  3. 使用离散优化算法来寻找一个近似的最优解,使得代价函数不断下降,从而不断调整网络节点间的权重;
  4. 代价函数收敛到一定程度以后结束迭代,保存节点间权重,此时神经网络训练完成;
  5. 使用该网络识别测试集中的手写数字样本,得出识别准确率。

经测试,该网络的识别准确率达到97.30%,与参考文献[3]中所得到的结果大致相当。下面列出被该神经网络识别错误的样本:

识别错误的数字

可以看出,这些样本自身确实具有一定的不规则性。考虑到本文并没有使用复杂的神经网络工具箱或者机器学习框架,所构建网络的拓扑结构较为简单,抽象与表述能力有限,因此得到这样的效果已经是较为理想的了。

图像预处理

当我们实现一个具有较高分类准确度的神经网络后,下一步的问题就是,如何将其应用到具体的使用场景中。换句话说,对于实际的包含手写数字的照片,我们如何将其中的数字一个个分离出来并加以识别?一种直观的想法是,对于所有像素点,我们在色彩空间中将其聚类,然后提取出与笔迹颜色(黑色)相近的像素点的集合,这些像素对应到图片上就是笔迹。但问题在于,笔迹像素仅仅占据整个图像的很小一部分,聚类算法很难单独将其分为一类,需要对聚类参数进行细致的微调才能获得较好的聚类效果,因此这种方法显然不具有通用性。

背景分离算法

我们考虑黑色笔迹的特点是,其RGB三个颜色分量基本都趋近于零。因此,如果我们将整个图像的三个颜色分量值求和,那么和最小的那些像素点,最有可能对应着笔迹。问题在于,我们如何确定一个阈值,使得低于这个阈值的像素点均为笔迹的像素,而高于这个阈值的像素点对应着图像背景呢?

首先让我们来观察一下整个过程。当我们动态变化这个阈值时,可以生成如下所示的gif动画:

动态过程示意

从上图中我们可以看出,在阈值非常低时,整个图像均为黑色;随着阈值的逐渐提升,慢慢有白色像素点开始出现在图像上,然后这些白色像素点开始互相连接形成数字;当阈值大到一定程度以后,白色像素开始连成一整片,并且最后占据了整个图像。

然后,我们再思考一下,在这个变化的过程中,图像中连通分量(8-连通)的数量发生了怎样的变化?当一开始整个图像为黑色时,图中连通分量数量为零;随着星星点点的白色像素的出现,连通分量数量开始增加。当阈值进一步升高时,这些白色像素点又会互相连接在一起形成数字,使得连通分量的数量减少;但随着阈值的继续升高,大量孤立的噪点开始出现,使得连通分量数目急剧增加。最后,所有白色噪点连成一片,连通分量数急剧下降至只有1个,也就是整个白色图像。这个过程大体如下图所示:

连通分量数随阈值的变化

最后我们考虑手写数字笔迹的特点:这些笔迹一定是连续的,也就是说,它们在图像上各自对应一个连通分量。因此,我们可以得到这样的结论,对于任何一张手写数字的图像,它一定具有这样的性质:随着阈值的升高,连通分量数一定是先从零开始增加,然后又减小,接下来再次急剧增加,最后急剧减小到1;在连通分量数的两次极大值之间一定存在一个极小值,这个极小值对应的连通分量数就是图中的手写数字个数。

综上所述,我们得到了一种自动化选取阈值的算法:只需要先求出整个动态变化过程中连通分量数随阈值的变化状况,然后取出两个极大值所对应的坐标区间,再在该区间中找到连通分量数最少的位置,该位置所对应的阈值即为我们所需要的值。这个阈值刚好能够使得所有笔迹都显现出来,并且不产生任何噪点,利用它我们可就以完美地将背景分离,并加强笔迹。

该算法的朴素Matlab实现如下:

function [ background ] = backgroundDetach( img )
    % Input: RGB image
    % Output: Enhanced binary image
    num_conn = zeros(1, 154);
    mask = repmat(sum(img, 3), [1,1,3]);
    start = -1;
    finish = -1;
    for i = 0:5:255*3
        im = img;
        im(mask > i) = 255;
        im = im2bw(im, min(i + 1, 255*3) / (255*3));
        im = ~im;
        [L, num] = bwlabel(im);
        idx = i / 5 + 1;
        num_conn(idx) = num;
        if (start == -1 && num > 0); start = idx; end;
        if (start ~= -1 && num == 1 ...
                && size(im, 1)*size(im, 2) == sum(sum(L==1)));
            finish = idx; break;
        end;
    end
    [peaks, location] = findpeaks(num_conn);
    [sorted, index] = sort(peaks, 'descend');
    min_pos = -1;
    min_val = 0;
    for i = location(index(2)):location(index(1))
        if (num_conn(i) <= min_val || min_pos == -1)
             min_pos = i;
             min_val = num_conn(i);
        end
    end;
    threshold = (min_pos - 1) * 5;
    img(mask > threshold) = 255;
    background = im2bw(img, (threshold + 1) / (255*3));
end

复杂度分析

上述实现利用了Matlab所提供的一个方便的求连通分量的函数,但其运行效率较低,因为每次改变阈值时都重新计算了整个图像的连通分量。下面我们分析该算法的理论复杂度,并给出一个高效率的实现方法。

首先我们对图像中每一点所对应的RGB颜色分量的值求和,相当于为每一个像素点赋予了一个位于区间[0, 3*255]的整数值。然后我们根据每个像素的值来对其进行分组,这个过程可以用O(MxN)的时空复杂度完成,其中M、N为图像的宽和高。接下来我们将图像置空,然后遍历[0, 3*255]这个区间,对于每一个值,找出其对应的所有像素,并将这些像素添加回图像中,同时确定连通分量数量的变化。这个过程中我们需要建立一个并查集数据结构,来维护连通分量间的关系。

在添加的过程中,我们检查当前像素周围邻域内是否存在已有像素,如果不存在,那么当前像素产生了一个新的连通分量,因此将现有连通分量个数加1,同时为这个像素赋予对应的连通分量标号;如果存在并且所有已有像素均属于同一个连通分量,那么当前像素所属的连通分量就是该连通分量;但如果几个已有像素属于不同的连通分量,那么就需要在并查集中对这些连通分量予以归并,然后再将归并后的连通分量编号赋予当前新像素。

该算法的伪代码示意图如下:

// 记录连通分量数量的变化趋势
all_conn[3*255 + 1] = {0}

for level <- [0, 3*255]
    for pixel_list <- get_pixels_by_level(level)
        for pixel in pixel_list
            if (pixel的8-邻域内存在像像素)
                if (pixel邻域内像素所属连通分量均为同一个,编号为idx)
                    image[pixel.row][pixel.col] = find_father(idx)
                else
                    在并查集中归并所有连通分量,得到编号为idx的分量
                    image[pixel.row][pixel.col] = find_father(idx)
                    num_conn对应减少
            else
                image[pixel.row][pixel.col] = new_conn_index()
    all_conn[level] = count_conn_num()

根据上述分析,显然可以得出结论:该算法的时间、空间复杂度均为O(MxN)

结果与总结

待识别的手写数字图像如下,是手写后使用手机相机拍摄的,未进行任何预处理:

待识别图像

利用上述提取背景的Matlab函数,结合适当的图像切分,我们得到的识别结果如下:

识别结果

从上图我们可以看出,对于自行构造的手写数字图片,本文所训练的神经网络的识别错误率稍高,10个数字错了3个。这一方面可能是由于图像的缩放使得数字产生了较大的形变,比如其中将“4”识别为“9”;另一方面也可能是因为我们所进行的图像预处理步骤,与MNIST数据集的预处理步骤存在一定差异所导致的。

参考文献

[1] MNIST数据集
[2] Coursera: Machine Learning
[3] Rahim M, Bengio Y, LeCun Y. Discriminative feature and model design for automatic speech recognition[C]//Eurospeech. 1997.

comments powered by Disqus
Published:
2015-07-10
分类:
Tag: