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

UNet继承tf.keras.Model训练报错:无_distribution_strategy属性

错误原因

你自定义的UNet类继承自tf.keras.Model,但存在两个核心初始化逻辑错误,导致父类内置属性没有被正确初始化:

  1. __init__方法没有调用父类tf.keras.Model的构造函数,父类内置的_distribution_strategy等属性完全没有被创建,因此调用compile时触发属性缺失报错。
  2. 错误重写了__call__方法,而非TensorFlow要求的call方法,覆盖了父类__call__中内置的模型初始化、输入校验等逻辑,进一步导致属性初始化不全。

修复代码

直接替换你的UNet类定义即可:

import tensorflow as tf

class UNet(tf.keras.Model):
    def __init__(self, img_shape=(256,256,256), num_class=1):
        # 新增:调用父类构造函数,完成内置属性初始化
        super(UNet, self).__init__()
        print ('build UNet ...')
        
        self.img_shape = img_shape+(1,)
        self.num_class = num_class

    def get_crop_shape(self, target, refer):
        # depth, the 4th dimension
        cd = (target.get_shape()[3] - refer.get_shape()[3])
        assert (cd >= 0)
        if cd % 2 != 0:
            cd1, cd2 = int(cd//2), int(cd//2) + 1
        else:
            cd1, cd2 = int(cd//2), int(cd//2)
        # width, the 3rd dimension
        cw = (target.get_shape()[2] - refer.get_shape()[2])
        assert (cw >= 0)
        if cw % 2 != 0:
            cw1, cw2 = int(cw//2), int(cw//2) + 1
        else:
            cw1, cw2 = int(cw//2), int(cw//2)
        # height, the 2nd dimension
        ch = (target.get_shape()[1] - refer.get_shape()[1])
        assert (ch >= 0)
        if ch % 2 != 0:
            ch1, ch2 = int(ch//2), int(ch//2) + 1
        else:
            ch1, ch2 = int(ch//2), int(ch//2)
    
        return (ch1, ch2), (cw1, cw2), (cd1, cd2)
    
    # 修改:方法名从__call__改为call,符合tf.keras.Model的正向传播定义规范
    def call(self, inputs):
        
        concat_axis = 4
        
        conv1 = tf.keras.layers.Conv3D(8, (3, 3, 3), activation='relu', padding='same', name='conv1_1')(inputs)
        conv1 = tf.keras.layers.Conv3D(8, (3, 3, 3), activation='relu', padding='same')(conv1)
        pool1 = tf.keras.layers.MaxPooling3D(pool_size=(2, 2, 2))(conv1)
        conv2 = tf.keras.layers.Conv3D(16, (3, 3, 3), activation='relu', padding='same')(pool1)
        conv2 = tf.keras.layers.Conv3D(16, (3, 3, 3), activation='relu', padding='same')(conv2)
        
        up_conv1 = tf.keras.layers.UpSampling3D(size=(2, 2, 2))(conv2)
        ch, cw, cd = self.get_crop_shape(conv1, up_conv1)
        crop_conv1 = tf.keras.layers.Cropping3D(cropping=(ch,cw,cd))(conv1)
        up1 = tf.keras.layers.concatenate([up_conv1, crop_conv1], axis=concat_axis)
        conv3 = tf.keras.layers.Conv3D(8, (3, 3, 3), activation='relu', padding='same')(up1)
        conv3 = tf.keras.layers.Conv3D(8, (3, 3, 3), activation='relu', padding='same')(conv3)
        
        ch, cw, cd = self.get_crop_shape(inputs, conv3)
        conv3 = tf.keras.layers.ZeroPadding3D(padding=((ch[0], ch[1]), (cw[0], cw[1]), (cd[0], cd[1])))(conv3)
        conv4 = tf.keras.layers.Conv3D(self.num_class, (1, 1, 1), activation="sigmoid")(conv3)
        
        return conv4

额外适配说明

你使用的TensorFlow 2.3版本较老,修复后如果出现模型权重保存失败的问题,可以调整ModelCheckpoint的文件名后缀,将.hdf5改为.h5,或者直接去掉后缀使用TensorFlow原生SavedModel格式保存,兼容性更好。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.07 09:36:04