对EfficientDet中加权融合方法的代码解读

本文主要解析了谷歌大脑EfficientDet中BiFPN的加权融合部分的代码,介绍了tensorflow的相关函数,如Variable、tf.cast()、tf.concat()、乘法运算以及reduce_sum()。代码实现了三种融合方法:attn(使用softmax),fastattn和sum(直接求和)。

谷歌大脑的《EfficientDet: Scalable and Efficient Object Detection》代码目前已经公布,代码链接:https://github.com/google/automl/tree/master/efficientdet
接下来将对BiFPN中加权融合部分进行解析。

对应代码入下:

 # Combine all nodes.
      dtype = nodes[0].dtype
      if config.weight_method == 'attn':
        edge_weights = [tf.cast(tf.Variable(1.0, name='WSM'), dtype=dtype)
                        for _ in range(len(fnode['inputs_offsets']))]
        normalized_weights = tf.nn.softmax(tf.stack(edge_weights))
        nodes = tf.stack(nodes, axis=-1)
        new_node = tf.reduce_sum(tf.multiply(nodes, normalized_weights), -1)
      elif config.weight_method == 'fastattn':
        edge_weights = [
            tf.nn.relu(tf.cast(tf.Variable(1.0, name='WSM'), dtype=dtype))
            for _ in range(len(fnode['inputs_offsets']))
        ]
        weights_sum = tf.add_n(edge_weights)
        nodes = [nodes[i] * edge_weights[i] / (weights_sum + 0.0001)
                 for i in range(len(nodes))]
        new_node = tf.add_n(nodes)
      elif config.weight_method == 'sum':
        new_node = tf.add_n(nodes)
      else:
        raise ValueError(
            'unknown weight_method {}'.format(config.weight_method))

代码是用tensorflow实现的,之前一直用的是pytorch,所以先对这部分代码中涉及到的一些函数进行解析。

W = tf.Variable(
                initial_value=tf.zeros([9, 5]),  
                        # 初始值,必填,张量或可以转换为张量的Python对象。初始值必须有指定一个形状,除非`validate_shape`设置为False。
                trainable=True,  
                        # 如果`True`,则默认值也将变量添加到图形中集合`GraphKeys.TRAINABLE_VARIABLES`。
                        #这个集合用作“Optimizer”类使用的默认变量列表
                collections=None,  
                        # 图表集合键的列表。新的变量被添加到这些集合。默认为`[GraphKeys.GLOBAL_VARIABLES]`。
                validate_shape=True, 
                        # 如果`False`,允许变量用初始化未知形状的值。如果“True”,默认的形状`initial_value`必须是已知的。
                 caching_device=None,  
                        # 可选设备字符串,描述变量的位置应该被缓存以供阅读。默认为变量的设备。如果不是“None”,则缓存在另一个设备上。
                        #典型的用途是缓存在使用变量 的Ops所在的设备上进行重复数据删除复制`Switch`和其他条件语句。
                 name='W',  
                        # 变量的可选名称。默认为“Variable”并获取自动去重(Variable_1,Variable_2....)。
                 variable_def=None,
                        # `VariableDef`协议缓冲区。如果不是“无”,则重新创建变量对象及其内容,引用变量的节点在图中,必须已经存在。
                        #图形没有改变。`variable_def`和其他参数是互斥的。
                dtype=tf.float32,
                        # 如果设置,initial_value将被转换为给定的类型。如果`None',数据类型将被保存
                        #(如果`initial_value`是一个张量),或者“convert_to_tensor”来决定。
                expected_shape=None,  
                        # 张量的Shape。如果设置,initial_value需要符合这个形状。
                 import_scope=None
                        # 可选的字符串。名称范围添加到`Variable.`仅在从协议缓冲区初始化时使用。
                    ) 

Vatiable是tensorflow的变量节点,通过Variable方法创建,并且需要传递初始值。在使用前需要通过tensorflow的初始化方法进行初始化。

cast(x, dtype, name=None)

tf.cast()函数的作用是执行 tensorflow 中张量数据类型转换,比如读入的图片如果是int8类型的,一般在要在训练前把图像的数据格式转换为float32。
第一个参数 x: 待转换的数据(张量)
第二个参数 dtype: 目标数据类型
第三个参数 name: 可选参数,定义操作的名称

tf.stack( values,axis=0)

将两个数组按照指定的方向进行叠加,生成一个新的数组。参数axis取0时表示按照x轴方向进行叠加,取1时表示按照y轴进行叠加。

tf.multiply(x,y,name=None)

乘法,位置相同的元素相乘。

tf.reduce_sum(
    input_tensor, 
    axis=None, 
    keepdims=None,
    name=None,
    reduction_indices=None, 
    keep_dims=None)

reduce_sum() 用于计算张量tensor沿着某一维度的和,可以在求和后降维。

  • input_tensor:待求和的tensor;
  • axis:指定的维,如果不指定,则计算所有元素的总和;
  • keepdims:是否保持原有张量的维度,设置为True,结果保持输入tensor的形状,设置为False,结果会降低维度,如果不传入这个参数,则系统默认为False;
  • name:操作的名称;
  • reduction_indices:在以前版本中用来指定轴,已弃用;
  • keep_dims:在以前版本中用来设置是否保持原张量的维度,已弃用;

总体来看这段代码,就是三种加权融合的方法:

  • attn
    在这里插入图片描述
    对每个 w i w_i wi使用一次softmax
  • fastattn
    在这里插入图片描述
  • sum
    直接求和
评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值