--- a/nntrainer/layers/split_layer.cpp +++ b/nntrainer/layers/split_layer.cpp @@ -161,12 +161,14 @@ // This is O(B * num_steps * split_number * split_w) vs the full // O(B * INIT_SEQ_LEN * split_number * split_w) of forwarding(). for (unsigned int b = 0; b < B; ++b) { - for (unsigned int s = 0; s < num_steps; ++s) { - for (unsigned int idx = 0; idx < split_number; ++idx) { - Tensor &output_ = context.getOutput(idx); - const float *src = input_.getAddress(b, 0, s, idx * split_w); - float *dst = output_.getAddress(b, 0, s, 0); - std::memcpy(dst, src, split_w * sizeof(float)); + for (unsigned int c = 0; c < input_.channel(); ++c) { + for (unsigned int s = 0; s < num_steps; ++s) { + for (unsigned int idx = 0; idx < split_number; ++idx) { + Tensor &output_ = context.getOutput(idx); + const float *src = input_.getAddress(b, c, s, idx * split_w); + float *dst = output_.getAddress(b, c, s, 0); + std::memcpy(dst, src, split_w * sizeof(float)); + } } } }