You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Keras实现求数字模7神经网络无效,准确率仅约1/7求排查

帮你搞定Keras模7预测的问题!

嘿,我来帮你排查这个问题!你的思路特别棒——用简单任务验证Keras流程,但几个关键的小错误导致网络根本没学到东西,结果就跟瞎蒙一样。咱们一步步拆解问题,再给你修正后的代码:

核心问题分析

1. 标签编码完全跑偏了

你现在把模7的结果(0-6之间的整数)转换成了20位二进制向量当标签,但categorical_crossentropy损失函数要的是one-hot编码的类别向量——也就是维度等于类别数(这里是7)的向量,每个位置对应一个类别的概率。比如模7结果是3,正确的标签应该是[0,0,0,1,0,0,0],而不是把3转成20位二进制。

2. 输出层和任务不匹配

你的输出层是20个单元的softmax,但我们的任务是分7类(0到6),所以输出层必须是7个单元,每个单元对应一个类别的概率。20个单元的softmax完全让网络摸不着头脑,不知道要预测什么。

3. 模型作用域有问题

你在main()里定义的model是局部变量,全局的predict()函数根本找不到它,运行的时候会直接报错。得调整一下变量作用域。

修正后的完整代码

import keras.models
import numpy as np
from python_toolbox import random_tools

RADIX = 7
model = None  # 全局声明模型变量

def _get_number(vector):
    return sum(x * 2 ** i for i, x in enumerate(vector))

def _get_mod_result(vector):
    return _get_number(vector) % RADIX

def _number_to_vector(number):
    binary_string = bin(number)[2:]
    if len(binary_string) > 20:
        raise NotImplementedError
    # 调整一下,返回(20,)的一维数组,方便后续处理
    bits = (((0,) * (20 - len(binary_string))) + tuple(map(int, binary_string)))[::-1]
    assert len(bits) == 20
    return np.array(bits)

def get_one_hot_label(vector):
    # 生成one-hot编码的标签,维度为7
    mod_result = _get_mod_result(vector)
    one_hot = np.zeros(RADIX)
    one_hot[mod_result] = 1
    return one_hot

def main():
    global model  # 声明使用全局模型变量
    model = keras.models.Sequential(
        (
            keras.layers.Dense(
                units=32,  # 稍微增加一点单元数,提升学习能力
                activation='relu',
                input_dim=20
            ),
            keras.layers.Dense(
                units=16,
                activation='relu'
            ),
            keras.layers.Dense(
                units=RADIX,  # 7个单元对应0-6这7个类别
                activation='softmax'
            )
        )
    )
    # 可以试试adam优化器,收敛比sgd快很多
    model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])
    
    # 生成训练数据
    data = np.random.randint(2, size=(10000, 20))
    # 生成one-hot标签
    labels = np.array([get_one_hot_label(row) for row in data])
    
    # 训练,增加epochs让网络学透
    model.fit(data, labels, epochs=20, batch_size=50, verbose=1)

def predict(number):
    vector = _number_to_vector(number)
    # 把(20,)的向量转成(1,20)的batch格式,符合模型输入要求
    vector = vector.reshape(1, 20)
    probabilities = model.predict(vector, verbose=0)
    # 取概率最大的索引,就是模7的结果
    return np.argmax(probabilities)

def is_correct_for_number(x):
    return predict(x) == x % RADIX

if __name__ == '__main__':
    main()
    # 测试部分
    print(f"预测7的模7结果:{predict(7)}(正确结果是0)")
    sample = random_tools.shuffled(range(2 ** 20))[:500]
    correct_count = sum(map(is_correct_for_number, sample))
    accuracy = correct_count / len(sample)
    print(f'Total accuracy: {accuracy:.4f}')
    print(f'(Accuracy of random algorithm is {1/RADIX:.2f})')

关键修改说明

  1. 标签系统重构:用get_one_hot_label生成7维one-hot标签,完全匹配categorical_crossentropy的要求。
  2. 输出层调整:改成7个单元的softmax,直接对应0-6的类别预测。
  3. 模型作用域修复:全局声明model,让predict能正常调用。
  4. 输入格式优化:_number_to_vector返回一维数组,predict里转成batch格式,符合Keras的输入规范。
  5. 优化器和网络结构微调:换成Adam优化器(收敛更快),调整中间层单元数,增加训练epochs,让网络充分学习。

运行这个修正后的代码,你会发现准确率能接近100%,完全不是之前的随机水平啦!

内容的提问来源于stack exchange,提问作者Ram Rachum

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.13 08:04:39