【问题标题】:Keras backend switch combined with tf.where not working as intendedKeras 后端开关结合 tf.where 无法按预期工作
【发布时间】:2022-11-23 00:02:46
【问题描述】:

我有一个自定义损失函数,我想将基于单热编码的值更改为特定范围内的值以计算 IOU。

这段代码的一部分是查看我在张量中有一个 1 的地方,否则张量为零。为此,我使用 tf.where 返回位置。我有一个形状为 [batch_size,S1,S2,12] 的向量,其中我只关心最后一个维度,这就是为什么我采用 tf.where 的 [...,2]。

现在经常发生我的预测全为零的情况,因为我的背景事件中没有任何值,而且我的网络会时不时地预测全零向量。这意味着 tf.where 将返回一个空张量。 这就是为什么我想使用 K.switch 检查张量是否为空,因为如果是,我希望返回零。

现在的问题是 K.switch 期望 then else 选项的形状具有相同的形状,但我需要我的输出具有形状 [batch_size,S1,S2,1]。我尝试过不同的东西,但我无法让它发挥作用。 我需要获得形状为 [batch_size,S1,S2,1] 的零点,或者我需要 where_box1 使 [batch_size,S1,S2,1] 带有浮点数。

现在实现的方式是,当 where_box1_temp 为空时,K.switch 返回一个空的零向量,这不是我想要的。 当我使用 tf.zeros([batch_size,S1,S2,1]) 时,它会抱怨 where_box1_temp 为空时条件形状不同....

where_box1_temp = tf.where(y_pred[...,C+1:C+13])[...,2]

where_box1 = K.switch(tf.equal(tf.size(where_box1_temp),0) , 
                          tf.zeros_like(where_box1_temp) , where_box1_temp)

【问题讨论】:

    标签: python tensorflow keras


    【解决方案1】:

    所以我找到了一个解决方法,也许这对其他人有帮助:

    where_box1_temp = tf.where(y_pred[...,C+1:C+13],[1,2,3,4,5,6,7,8,9,10,11,12],0)
    
    where_box1 = tf.reshape(K.sum(where_box1_temp,axis=3),[batch_size,5,5])
    

    这使我能够获得所需形状的张量,其中所有背景/零预测值均为 0,而无需使用 k.switch 并且不会遇到任何空维度或类似问题。

    【讨论】:

      猜你喜欢
      • 2017-09-22
      • 2016-10-13
      • 2021-11-13
      • 2016-12-26
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2019-02-19
      • 2021-10-02
      相关资源
      最近更新 更多