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

Method setup_data

src/layers/learning/base_convolution.cpp:328–425  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

326
327template <typename TensorDataType, El::Device Device>
328void base_convolution_layer<TensorDataType, Device>::setup_data(
329 size_t max_mini_batch_size)
330{
331 data_type_layer<TensorDataType>::setup_data(max_mini_batch_size);
332
333 // Tensor dimensions
334 const auto& input_dims = this->get_input_dims();
335 const auto& output_dims = this->get_output_dims();
336 const auto& kernel_dims_ = this->get_kernel_dims();
337 std::vector<size_t> kernel_dims(kernel_dims_.begin(), kernel_dims_.end());
338 const auto kernel_size = get_linear_size(kernel_dims);
339
340 // Initialize default weights if none are provided
341 if (this->num_weights() > 2) {
342 LBANN_ERROR("attempted to setup layer \"",
343 this->get_name(),
344 "\" "
345 "with an invalid number of weights "
346 "(expected at most 2, "
347 "found ",
348 this->num_weights(),
349 ")");
350 }
351 if (m_bias_scaling_factor != El::TypeTraits<ScalingType>::Zero()) {
352 this->set_num_weights(2);
353 }
354 else {
355 this->set_num_weights(1);
356 }
357 if (!this->has_weights(0)) {
358 auto w = std::make_shared<WeightsType>(*this->get_comm());
359 auto init = std::make_unique<he_initializer<TensorDataType>>(
360 probability_distribution::gaussian);
361 auto opt = this->m_model->template create_optimizer<TensorDataType>();
362
363 w->set_name(this->get_name() + "_kernel");
364 w->set_initializer(std::move(init));
365 w->set_optimizer(std::move(opt));
366 this->set_weights(0, w);
367 this->m_model->add_weights(std::move(w));
368 }
369 auto& kernel_weights = this->get_weights(0);
370
371 // Initialize variance scaling initialization
372 if (auto* initializer = kernel_weights.get_initializer()) {
373 set_fan_in(*initializer, kernel_size / output_dims[0]);
374 set_fan_out(*initializer, kernel_size / input_dims[0]);
375 }
376
377 // Initialize weight matrices
378 auto dist = this->get_prev_activations().DistData();
379 dist.colDist = El::STAR;
380 dist.rowDist = El::STAR;
381 kernel_weights.set_dims(kernel_dims);
382 kernel_weights.set_matrix_distribution(dist);
383
384 // Set up bias if needed.
385 if (m_bias_scaling_factor != El::TypeTraits<ScalingType>::Zero()) {

Callers

nothing calls this directly

Calls 15

get_linear_sizeFunction · 0.85
ZeroClass · 0.85
set_fan_inFunction · 0.85
set_fan_outFunction · 0.85
get_input_dimsMethod · 0.80
get_output_dimsMethod · 0.80
beginMethod · 0.80
endMethod · 0.80
num_weightsMethod · 0.80
set_num_weightsMethod · 0.80
has_weightsMethod · 0.80
set_initializerMethod · 0.80

Tested by

no test coverage detected