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

TensorFlow Keras国际象棋ANN自定义损失函数梯度缺失问题求助

问题解答

1. 损失函数应返回什么值?

Keras自定义损失函数必须返回TensorFlow张量(Tensor),不能是Python列表或NumPy数组。具体要求:

  • 可返回与输入批次同维度的损失张量(每个样本对应一个损失值),Keras会自动计算批次平均值用于反向传播。
  • 也可直接返回一个标量张量(整个批次的平均损失)。
  • 核心禁忌:不能使用.numpy()将张量转为NumPy数组,也不能用np.argmax()这类脱离TensorFlow计算图的操作——这些操作会断开梯度传播链路,导致模型无法计算参数梯度,出现你遇到的"No gradients provided"错误,必须用TensorFlow内置函数(如tf.argmax()、tf.lookup)替代。

2. 计算出的损失如何用于ANN训练?

损失是模型训练的核心优化目标,流程如下:

  1. 损失函数计算当前模型预测与真实标签之间的差异(包含非法走法的惩罚)。
  2. Keras自动通过反向传播算法,计算损失相对于模型所有可训练参数(权重、偏置)的梯度。
  3. 优化器(如Adam、SGD)根据梯度更新参数,逐步减小损失,让模型的预测越来越接近合法且正确的走法。
  • 只有当损失计算全程在TensorFlow计算图中完成(无断开操作),梯度才能正常传播,参数更新才能生效。

3. 其他抑制非法走法的方法?

除了在损失函数中加惩罚,还有更高效的方案:

  • 输出层掩码法:提前生成当前棋盘的所有合法走法掩码(一个与输出维度相同的张量,合法位置为1,非法为0),将模型输出y_pred与掩码相乘,再计算交叉熵损失。这样模型只能从合法走法中选择,无需在损失里判断合法性。
  • 数据过滤:训练前直接过滤掉非法走法的样本,只让模型学习合法走法的标签,从源头上减少非法走法的预测概率。
  • 辅助正则项:训练一个辅助模型判断走法合法性,将辅助模型的输出作为主模型的正则损失项,让主模型同时优化预测准确性和合法性。
  • 强化学习(RL):参考AlphaZero思路,让模型通过自我对弈学习,环境会自动惩罚非法走法(如直接判负),模型会在试错中逐渐学会规避非法走法。

修正后的损失函数核心示例

import tensorflow as tf

def source_loss(y_true, y_pred):
    # 转换为TensorFlow张量操作,保持计算图连续性
    source_tile_true = tf.cast(y_true[:, 0, 0], tf.int32)
    target_tile_true = tf.cast(y_true[:, 1, 0], tf.int32)
    game_board = tf.cast(y_true[:, 2], tf.int32)
    
    # 用tf.argmax替代np.argmax
    source_tile_pred = tf.argmax(y_pred, axis=1, output_type=tf.int32)
    
    # 基础损失用稀疏分类交叉熵,更适合分类任务
    base_loss = tf.keras.losses.sparse_categorical_crossentropy(
        tf.expand_dims(source_tile_true, axis=1), y_pred
    )
    
    # 构建棋子类型映射表(TensorFlow版)
    keys = tf.constant([0,1,2,3,4,5,6,7,8,9,10,11,12], dtype=tf.int32)
    values = tf.constant(['_','p','r','n','b','q','k','P','R','N','B','Q','K'], dtype=tf.string)
    int_to_piece_table = tf.lookup.StaticHashTable(
        tf.lookup.KeyValueTensorInitializer(keys, values), default_value='_'
    )
    piece_type = int_to_piece_table.lookup(tf.gather(game_board, source_tile_true, batch_dims=1))
    
    # 黑兵合法走法的TensorFlow实现
    def is_valid_move_black_pawn(source, target):
        source_coord = tf.stack([source // 8, source % 8], axis=1)
        target_coord = tf.stack([target // 8, target % 8], axis=1)
        return tf.logical_or(
            tf.equal(target_coord - source_coord, tf.constant([1, 0])),
            tf.logical_and(tf.equal(source_coord[:,0], 1), tf.equal(target_coord - source_coord, tf.constant([2, 0])))
        )
    
    # 白兵合法走法的TensorFlow实现
    def is_valid_move_white_pawn(source, target):
        source_coord = tf.stack([source // 8, source % 8], axis=1)
        target_coord = tf.stack([target // 8, target % 8], axis=1)
        return tf.logical_or(
            tf.equal(target_coord - source_coord, tf.constant([-1, 0])),
            tf.logical_and(tf.equal(source_coord[:,0], 6), tf.equal(target_coord - source_coord, tf.constant([-2, 0])))
        )
    
    # 判断走法合法性
    valid_mask = tf.case([
        (tf.equal(piece_type, 'p'), lambda: is_valid_move_black_pawn(source_tile_pred, target_tile_true)),
        (tf.equal(piece_type, 'P'), lambda: is_valid_move_white_pawn(source_tile_pred, target_tile_true)),
        # 其他棋子类型的判断逻辑同理补充
    ], default=lambda: tf.constant(False, dtype=tf.bool))
    
    # 添加非法走法惩罚
    penalty = tf.where(tf.logical_not(valid_mask), 100.0, 0.0)
    total_loss = base_loss + penalty
    
    return total_loss

内容的提问来源于stack exchange,提问作者Oleksandr Nekrashevych

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 00:54:59