使用TransUNet库时model.parameters()无法调用,寻求替代方案
解决TransUNet模型无法使用PyTorch优化器的问题
问题根源
你使用的TransUNet-tf是TensorFlow/Keras版本的实现,但你却尝试调用PyTorch的torch.optim.Adam优化器,两个框架的API完全不兼容,因此PyTorch的parameters()等方法对这个TensorFlow模型无效。
可行解决方案
方案1:适配TensorFlow/Keras生态使用
既然用了TensorFlow版本的TransUNet,就改用TensorFlow的优化器:
import tensorflow as tf from transunet import TransUNet # 初始化模型 model = TransUNet(image_size=224, pretrain=True) # 编译模型时指定TensorFlow优化器 model.compile( optimizer=tf.keras.optimizers.Adam(learning_rate=1e-4), loss=tf.keras.losses.BinaryCrossentropy(), # 根据你的任务选择合适损失函数 metrics=['accuracy'] )
方案2:切换到PyTorch版本的TransUNet
如果一定要用PyTorch的优化器,需要更换为PyTorch实现的TransUNet仓库,之后就可以正常使用torch.optim.Adam(model.parameters(), lr=1e-4)来创建优化器。
内容的提问来源于stack exchange,提问作者Ankit Wankhede
相关产品推荐
相关产品推荐

