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

TensorFlow相同代码在Colab与本地环境运行结果不一致的技术咨询

我来帮你分析一下问题的核心原因,然后给出针对性的解决方案:

问题原因拆解

1. 形状不匹配引发的计算异常

你自定义的CategoricalTruePositives Metric中,把y_pred reshaped成了(-1,1)的二维张量,但y_true是(-1,)的一维张量。虽然TensorFlow支持广播机制,但在Colab的TensorFlow 2.4.1环境中,这种跨维度的比较可能触发了未预期的计算逻辑,导致累加的真正例数只有预期的一半左右。而本地环境的TensorFlow实现对这种广播的处理更符合预期,所以得到了正确的结果。

2. 浮点数精度导致的小数问题

默认情况下,Metric的权重变量使用float32类型,在多次累加小数值(比如每个batch的正确数)时,会出现精度损失,最终结果带有小数。而你的本地环境可能因为配置差异(比如默认使用float64),累加精度更高,所以结果是整数。


解决方案

修正自定义Metric代码

调整update_state方法,确保y_true和y_pred的形状完全一致,同时指定更高精度的float64类型来避免小数问题:

class CategoricalTruePositives(keras.metrics.Metric):
    def __init__(self, name="categorical_true_positives", **kwargs):
        super(CategoricalTruePositives, self).__init__(name=name, **kwargs)
        # 使用float64类型存储累加值,避免精度损失
        self.true_positives = self.add_weight(name="ctp", initializer="zeros", dtype=tf.float64)

    def update_state(self, y_true, y_pred, sample_weight=None):
        # 获取每个样本的预测类别,形状为(batch_size,)
        y_pred = tf.argmax(y_pred, axis=1)
        # 确保y_true的形状与y_pred完全一致(一维)
        y_true = tf.reshape(y_true, shape=(-1,))
        # 直接比较预测标签与真实标签
        values = tf.cast(tf.equal(y_true, y_pred), tf.float64)
        
        if sample_weight is not None:
            sample_weight = tf.cast(sample_weight, tf.float64)
            values = tf.multiply(values, sample_weight)
        
        self.true_positives.assign_add(tf.reduce_sum(values))

    def result(self):
        return self.true_positives

    def reset_states(self):
        self.true_positives.assign(0.0)

额外建议:统一训练随机性

为了让Colab和本地的训练结果完全一致,建议在代码开头设置随机种子,消除初始化和训练过程中的随机性差异:

import tensorflow as tf
import numpy as np

# 设置全局随机种子
tf.random.set_seed(42)
np.random.seed(42)

修改后,你在两个环境中得到的categorical_true_positives数值会与acc * 样本数完全匹配,且结果为整数,同时训练的loss和acc也会基本一致。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.29 16:13:13