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

将TensorFlow2.13自定义Metric适配2.2,解决张量布尔转换错误

TensorFlow 2.2适配自定义Distance Metric的解决方案

问题背景

环境:Ubuntu 20.04、Python 3.8
原有代码基于TensorFlow 2.13.0开发,包含U-Net和自定义损失函数,需适配仅支持TensorFlow 2.2.0(或更低版本)的GPU集群。核心问题出在自定义Distance Metric的update_state方法:

  1. 直接用Python if判断tf.Tensor的布尔值,触发OperatorNotAllowedInGraphError
  2. 尝试用@tf.function封装逻辑后,又因切片赋值的变量操作限制触发ValueError

原代码

class Distance(tf.keras.metrics.Metric):


    def __init__(self, name='DistanceMetric', distance='cm', sigma=2.5, data_size=None,
                 validation_size=None, points=None, point=None, percentile=None):
        super(Distance, self).__init__(name=name)
        self.counter = tf.Variable(initial_value=0, dtype=tf.int32)
        self.distance = distance
        self.sigma = sigma
        self.percentile = percentile
        if percentile is not None and point is not None:
            assert (type(percentile) == float)
            self.percentile_idx = tf.Variable(tf.cast(tf.round(percentile * validation_size), dtype=tf.int32))
        else:
            self.percentile_idx = None
        self.point = point
        self.points = points
        self.cache = tf.Variable(initial_value=tf.zeros([validation_size, points]),
                                 shape=[validation_size, points])
        self.val_size = validation_size

    def update_state(self, y_true, y_pred, sample_weight=None):
        n, h, w, p = tf.shape(y_pred)[0], tf.shape(y_pred)[1], tf.shape(y_pred)[2], tf.shape(y_pred)[3]
        y_true = normal_distribution(self.sigma, y_true[:, :, 0], y_true[:, :, 1], h=h, w=w, n=n, p=p)
        if self.distance == 'cm':
            x1, y1 = cm(y_true, h=h, w=w, n=n, p=p)
            x2, y2 = cm(y_pred, h=h, w=w, n=n, p=p)
            d = ((x1 - x2) ** 2 + (y1 - y2) ** 2) ** 0.5
            d = d[:, :, 0]
        elif self.distance == 'argmax':
            d = (tf.cast(tf.reduce_sum(((argmax_2d(y_true) - argmax_2d(y_pred)) ** 2), axis=1),
                         dtype=tf.float32)) ** 0.5

        temp = tf.minimum(self.counter + n, self.val_size)
        if self.counter <= self.val_size:
            self.cache[self.counter:temp, :].assign(d[0:(temp-self.counter), :])

        self.counter.assign(self.counter + n)

    def result(self):
        if self.percentile_idx is not None:
            temp = tf.sort(self.cache[:self.val_size, self.point], axis=0, direction='ASCENDING')
            return temp[self.percentile_idx]
        elif self.point is not None:
            return tf.reduce_mean(self.cache[:, self.point], axis=0)
        else:
            return tf.reduce_mean(self.cache, axis=None)

    def reset_states(self):
        self.cache.assign(tf.zeros_like(self.cache))
        self.counter.assign(0)
        if self.percentile is not None and self.point is not None:
            self.percentile_idx.assign(tf.cast(self.val_size * self.percentile, dtype=tf.int32))

首次报错信息

/trinity/home/r084755/DRF_AI/distal-radius-fractures-x-pa-and-lateral-to-clinic/Code files/LandmarkDetection.py:144 update_state
        if tf.math.less_equal(self.counter, self.val_size):         # Updated from self.counter <= self.val_size:
    /opt/ohpc/pub/easybuild/software/TensorFlow/2.2.0-fosscuda-2019b-Python-3.7.4/lib/python3.7/site-packages/tensorflow/python/framework/ops.py:778 __bool__
        self._disallow_bool_casting()
    /opt/ohpc/pub/easybuild/software/TensorFlow/2.2.0-fosscuda-2019b-Python-3.7.4/lib/python3.7/site-packages/tensorflow/python/framework/ops.py:545 _disallow_bool_casting
        "using a `tf.Tensor` as a Python `bool`")
    /opt/ohpc/pub/easybuild/software/TensorFlow/2.2.0-fosscuda-2019b-Python-3.7.4/lib/python3.7/site-packages/tensorflow/python/framework/ops.py:532 _disallow_when_autograph_enabled
        " decorating it directly with @tf.function.".format(task))

OperatorNotAllowedInGraphError: using a `tf.Tensor` as a Python `bool` is not allowed: AutoGraph did not convert this function. Try decorating it directly with @tf.function.

尝试修改后的代码及报错

修改代码片段:

@tf.function
def myfunc(counter, val_size, cache):
    temp = tf.minimum(counter + n, val_size-1)
    if counter <= val_size:
        return cache[counter:temp, :].assign(d[0:(temp-counter), :])
    return cache

self.cache = myfunc(self.counter, self.val_size, self.cache)
self.counter.assign(self.counter + n)

报错信息:

/opt/ohpc/pub/easybuild/software/TensorFlow/2.2.0-fosscuda-2019b-Python-3.7.4/lib/python3.7/site-packages/tensorflow/python/keras/engine/training.py:571 train_function  *
    outputs = self.distribute_strategy.run(
/trinity/home/r084755/DRF_AI/distal-radius-fractures-x-pa-and-lateral-to-clinic/Code files/LandmarkDetection.py:158 myfunc  *
    return cache[counter:temp, :].assign(d[0:(temp-counter), :])
/opt/ohpc/pub/easybuild/software/TensorFlow/2.2.0-fosscuda-2019b-Python-3.7.4/lib/python3.7/site-packages/tensorflow/python/ops/array_ops.py:1160 assign  **
    raise ValueError("Sliced assignment is only supported for variables")

ValueError: Sliced assignment is only supported for variables

适配TF2.2的修改方案

核心修改点:

  1. 用TensorFlow原生控制流tf.cond替代Python if,避免将tf.Tensor作为Python布尔值判断
  2. 直接对tf.Variable执行切片赋值,不在函数中返回赋值结果

修改后的完整Distance类:

class Distance(tf.keras.metrics.Metric):


    def __init__(self, name='DistanceMetric', distance='cm', sigma=2.5, data_size=None,
                 validation_size=None, points=None, point=None, percentile=None):
        super(Distance, self).__init__(name=name)
        self.counter = tf.Variable(initial_value=0, dtype=tf.int32)
        self.distance = distance
        self.sigma = sigma
        self.percentile = percentile
        if percentile is not None and point is not None:
            assert (type(percentile) == float)
            self.percentile_idx = tf.Variable(tf.cast(tf.round(percentile * validation_size), dtype=tf.int32))
        else:
            self.percentile_idx = None
        self.point = point
        self.points = points
        self.cache = tf.Variable(initial_value=tf.zeros([validation_size, points]),
                                 shape=[validation_size, points])
        self.val_size = validation_size

    def update_state(self, y_true, y_pred, sample_weight=None):
        n, h, w, p = tf.shape(y_pred)[0], tf.shape(y_pred)[1], tf.shape(y_pred)[2], tf.shape(y_pred)[3]
        y_true = normal_distribution(self.sigma, y_true[:, :, 0], y_true[:, :, 1], h=h, w=w, n=n, p=p)
        if self.distance == 'cm':
            x1, y1 = cm(y_true, h=h, w=w, n=n, p=p)
            x2, y2 = cm(y_pred, h=h, w=w, n=n, p=p)
            d = ((x1 - x2) ** 2 + (y1 - y2) ** 2) ** 0.5
            d = d[:, :, 0]
        elif self.distance == 'argmax':
            d = (tf.cast(tf.reduce_sum(((argmax_2d(y_true) - argmax_2d(y_pred)) ** 2), axis=1),
                         dtype=tf.float32)) ** 0.5

        temp = tf.minimum(self.counter + n, self.val_size)
        
        # 定义两个分支函数,tf.cond要求分支返回相同类型结果
        def assign_cache():
            # 直接对self.cache变量执行切片赋值
            self.cache[self.counter:temp, :].assign(d[0:(temp - self.counter), :])
            return self.cache
        
        def do_nothing():
            return self.cache
        
        # 用tf.cond替代Python if,符合TF图模式要求
        tf.cond(tf.less_equal(self.counter, self.val_size), assign_cache, do_nothing)

        self.counter.assign(self.counter + n)

    def result(self):
        if self.percentile_idx is not None:
            temp = tf.sort(self.cache[:self.val_size, self.point], axis=0, direction='ASCENDING')
            return temp[self.percentile_idx]
        elif self.point is not None:
            return tf.reduce_mean(self.cache[:, self.point], axis=0)
        else:
            return tf.reduce_mean(self.cache, axis=None)

    def reset_states(self):
        self.cache.assign(tf.zeros_like(self.cache))
        self.counter.assign(0)
        if self.percentile is not None and self.point is not None:
            self.percentile_idx.assign(tf.cast(self.val_size * self.percentile, dtype=tf.int32))

修改说明

  • tf.cond是TensorFlow图模式下的条件控制流,可接收tf.Tensor类型的判断条件,避免Python控制流的类型错误
  • 切片赋值必须直接操作类内的tf.Variable(self.cache),不能将赋值结果返回后重新赋值给变量,否则会触发“仅支持变量的切片赋值”错误
  • 保留原有逻辑的同时,完全适配TF2.2的图模式执行规则

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 02:54:50