MCPcopy Create free account
hub / github.com/LBANN/lbann / setup_data

Method setup_data

src/layers/learning/fully_connected.cpp:120–247  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

118
119template <typename TensorDataType, data_layout T_layout, El::Device Dev>
120void fully_connected_layer<TensorDataType, T_layout, Dev>::setup_data(
121 size_t max_mini_batch_size)
122{
123 data_type_layer<TensorDataType>::setup_data(max_mini_batch_size);
124
125 // Initialize default weights if none are provided
126 if (this->num_weights() > 2) {
127 LBANN_ERROR("attempted to setup ",
128 this->get_name(),
129 " with an invalid number of weights");
130 }
131 if (m_bias_scaling_factor != El::TypeTraits<TensorDataType>::Zero()) {
132 this->set_num_weights(2);
133 }
134 else {
135 this->set_num_weights(1);
136 }
137 if (!this->has_weights(0)) {
138 auto w = std::make_shared<WeightsType>(*this->get_comm());
139 auto init = std::make_unique<he_initializer<TensorDataType>>(
140 probability_distribution::gaussian);
141 auto opt = this->m_model->template create_optimizer<TensorDataType>();
142 w->set_name(this->get_name() + "_linearity_weights");
143 w->set_initializer(std::move(init));
144 w->set_optimizer(std::move(opt));
145 this->set_weights(0, w);
146 this->m_model->add_weights(std::move(w));
147 }
148 auto& linearity_weights = this->get_weights(0);
149
150 // Initialize variance scaling initialization
151 if (auto* initializer = linearity_weights.get_initializer()) {
152 set_fan_in(*initializer, this->get_input_size());
153 set_fan_out(*initializer, this->get_output_size());
154 }
155
156 // Input and output dimensions
157 const auto& input_dims_ = this->get_input_dims();
158 const auto& output_dims_ = this->get_output_dims();
159 std::vector<size_t> input_dims(input_dims_.begin(), input_dims_.end());
160 std::vector<size_t> output_dims(output_dims_.begin(), output_dims_.end());
161
162 // Setup linearity weights
163 auto linearity_dist = this->get_prev_activations().DistData();
164 if (linearity_dist.colDist != El::MC || linearity_dist.rowDist != El::MR) {
165 linearity_dist.colDist = El::STAR;
166 linearity_dist.rowDist = El::STAR;
167 }
168 if (m_transpose) {
169 linearity_weights.set_dims(input_dims, output_dims);
170 }
171 else {
172 linearity_weights.set_dims(output_dims, input_dims);
173 }
174 linearity_weights.set_matrix_distribution(linearity_dist);
175
176 // Set up bias if needed.
177 if (m_bias_scaling_factor != El::TypeTraits<TensorDataType>::Zero()) {

Callers

nothing calls this directly

Calls 15

ZeroClass · 0.85
set_fan_inFunction · 0.85
set_fan_outFunction · 0.85
num_weightsMethod · 0.80
set_num_weightsMethod · 0.80
has_weightsMethod · 0.80
set_initializerMethod · 0.80
set_optimizerMethod · 0.80
get_initializerMethod · 0.80
get_input_dimsMethod · 0.80
get_output_dimsMethod · 0.80
beginMethod · 0.80

Tested by

no test coverage detected