【问题标题】:Is there a way to know how many parameters does an object detection model have, in tensorflow object detection API?在 tensorflow 对象检测 API 中,有没有办法知道对象检测模型有多少参数?
【发布时间】:2019-11-12 02:23:17
【问题描述】:

我使用张量对象检测 (TFOD) API 训练不同的模型,我想知道为给定模型训练了多少参数。

我运行更快的 RCNN、SSD、RFCN 以及不同的图像分辨率,我想知道训练了多少参数。有没有办法做到这一点?

我已经尝试在这里找到答案 How to count total number of trainable parameters in a tensorflow model? 没有运气。

这是我在model_main.py的第103行添加的代码:

print("Training {} parameters".format(np.sum([np.prod(v.get_shape().as_list()) for v in tf.trainable_variables()]))

我认为问题在于我没有访问 TFOD 正在运行的 tf.Session(),因此我的代码总是返回 0.0 个参数(尽管训练策略很好,希望能训练数百万个参数),但我没有不知道怎么解决这个问题。

【问题讨论】:

  • 训练多少参数是什么意思?您的意思是您的模型中有多少个参数?
  • @BlueRineS 是的,我需要知道模型中有多少参数
  • 嘿伙计,如果您可以发布一些参数,例如 ssd、更快的 rcnn 等,那将非常有用。

标签: python tensorflow object-detection-api


【解决方案1】:

使用 export_inference_graph.py 时,脚本还会分析您的模型,并计算参数和 FLOPS(如果可能)。 如果看起来像这样:

_TFProfRoot (--/# total params)
  FeatureExtractor (--/# params)
  ...
  WeightSharedConvolutionalBoxPredictor (--/# params)
  ...

【讨论】:

    【解决方案2】:

    TFOD API 使用tf.estimator.Estimator 进行训练和评估。 Estimator 对象提供了获取所有变量的函数,Estimator.get_variable_names() (reference)。

    您可以在estimator.train_and_evaluate() (here) 之后添加此行print(estimator.get_variable_names())

    训练完成后,您将看到打印的所有变量名称。要更快地查看结果,您只需训练 1 步即可。

    【讨论】:

    • 这是您正在寻找的答案吗?还是我回答错了?
    猜你喜欢
    • 2018-01-17
    • 2019-11-09
    • 1970-01-01
    • 1970-01-01
    • 2018-12-26
    • 2020-11-02
    • 1970-01-01
    • 2018-03-30
    • 2017-12-02
    相关资源
    最近更新 更多