[ Web Proxy ]
URL:
Viewing: https://raw.githubusercontent.com/mirsys/CaffeOnACL/master/unit_tests/test_softmax_layer.cpp [Back]  [Original]

#include 
#include 

#include "gtest/gtest.h"

#include "caffe/blob.hpp"
#include "caffe/common.hpp"
#include "caffe/filler.hpp"
#include "caffe/layers/softmax_layer.hpp"

#ifdef USE_CUDNN
#include "caffe/layers/cudnn_softmax_layer.hpp"
#endif

#include "caffe/test/test_caffe_main.hpp"
#include "caffe/test/test_gradient_check_util.hpp"

namespace caffe {

template 
class SoftmaxLayerTest : public MultiDeviceTest {
  typedef typename TypeParam::Dtype Dtype;
 protected:
  SoftmaxLayerTest()
      : blob_bottom_(new Blob(2, 10, 1, 1)),
        blob_top_(new Blob()) {
    // fill the values
    FillerParameter filler_param;
    GaussianFiller filler(filler_param);
    filler.Fill(this->blob_bottom_);
    blob_bottom_vec_.push_back(blob_bottom_);
    blob_top_vec_.push_back(blob_top_);
  }
  virtual ~SoftmaxLayerTest() { delete blob_bottom_; delete blob_top_; }
  Blob* const blob_bottom_;
  Blob* const blob_top_;
  vector blob_bottom_vec_;
  vector blob_top_vec_;
};


typedef ::testing::Types float_only;

#define TestDtypesAndDevices float_only


TYPED_TEST_CASE(SoftmaxLayerTest, TestDtypesAndDevices);

TYPED_TEST(SoftmaxLayerTest, TestForward) {
  typedef typename TypeParam::Dtype Dtype;
  LayerParameter layer_param;


 layer_param.set_type("Softmax");

  shared_ptr new_layer=
    LayerRegistry::CreateLayer(layer_param);

  shared_ptr layer=
   boost::static_pointer_cast (new_layer);

//  layer=shared_ptr(new  SoftmaxLayer(layer_param));

  layer->SetUp(this->blob_bottom_vec_, this->blob_top_vec_);
  layer->Forward(this->blob_bottom_vec_, this->blob_top_vec_);


  // Test sum
  for (int i = 0; i < this->blob_bottom_->num(); ++i) {
    for (int k = 0; k < this->blob_bottom_->height(); ++k) {
      for (int l = 0; l < this->blob_bottom_->width(); ++l) {
        Dtype sum = 0;
        for (int j = 0; j < this->blob_top_->channels(); ++j) {
          sum += this->blob_top_->data_at(i, j, k, l);
        }
        EXPECT_GE(sum, 0.999);
        EXPECT_LE(sum, 1.001);
        // Test exact values
        Dtype scale = 0;
        for (int j = 0; j < this->blob_bottom_->channels(); ++j) {
          scale += exp(this->blob_bottom_->data_at(i, j, k, l));
        }
        for (int j = 0; j < this->blob_bottom_->channels(); ++j) {
          EXPECT_GE(this->blob_top_->data_at(i, j, k, l) + 1e-4,
              exp(this->blob_bottom_->data_at(i, j, k, l)) / scale)
              data_at(i, j, k, l)) / scale)
              

Web Proxy Viewer  |  New URL  |  Original Page