如何为tf.keras.layers.RandomFlip设置随机翻转的触发概率?
tf.keras.layers.RandomFlip 触发概率设置说明
可以设置随机翻转的触发概率,具体使用规则和注意事项如下:
- 官方提供了参数
p用来指定触发概率,取值范围为0到1之间的浮点数,数值代表翻转操作被执行的概率。例如p=0.5就表示每次调用该层时,有50%的概率执行指定方向的翻转,50%的概率直接返回原输入。 - 该参数是TensorFlow 2.10及以上版本新增的配置项,同系列的其他随机数据增强层(比如
RandomRotation、RandomZoom、RandomBrightness等)都支持该参数自定义触发概率。 - 如果使用的是TensorFlow 2.9及更早版本,直接使用
p参数会触发参数不存在的报错,此时要么升级TensorFlow版本到2.10以上,要么自定义包装层实现概率控制逻辑。
你提供的示例写法符合高版本TensorFlow的使用规范,可直接运行:
import tensorflow as tf from tensorflow import keras def augmentation(): data_augmentation = keras.Sequential([ keras.layers.RandomFlip("horizontal", p=0.5), keras.layers.RandomRotation(0.2, p=0.5) ]) return data_augmentation
如果使用旧版本TensorFlow,可以参考如下自定义层实现同等效果:
class RandomFlipWithProb(keras.layers.Layer): def __init__(self, flip_mode="horizontal", prob=0.5, **kwargs): super().__init__(**kwargs) self.flip_layer = keras.layers.RandomFlip(flip_mode) self.prob = prob def call(self, inputs, training=None): if not training: return inputs if tf.random.uniform(shape=[]) < self.prob: return self.flip_layer(inputs) return inputs # 使用自定义层构造数据增强流水线 def augmentation(): data_augmentation = keras.Sequential([ RandomFlipWithProb("horizontal", prob=0.5), keras.layers.RandomRotation(0.2) ]) return data_augmentation
内容的提问来源于stack exchange,提问作者haofeng
相关产品推荐
相关产品推荐

