169 if (!model.CheckIfTensorAlreadyExist(fNX)) {
170 throw std::runtime_error(
"TMVA SOFIE GRU Op input tensor " + fNX +
" is not found in model.");
172 fShapeX = model.GetTensorShape(fNX);
173 if (fShapeX.size() != 3) {
174 throw std::runtime_error(
"TMVA SOFIE GRU Op input tensor " + fNX +
" is not of 3 dimensions.");
176 if (!model.CheckIfTensorAlreadyExist(fNW)) {
177 throw std::runtime_error(
"TMVA SOFIE GRU Op input tensor " + fNW +
" is not found in model.");
179 fShapeW = model.GetTensorShape(fNW);
180 if (fShapeW.size() != 3) {
181 throw std::runtime_error(
"TMVA SOFIE GRU Op input tensor " + fNW +
" is not of 3 dimensions.");
183 if (!model.CheckIfTensorAlreadyExist(fNR)) {
184 throw std::runtime_error(
"TMVA SOFIE GRU Op input tensor " + fNR +
" is not found in model.");
186 fShapeR = model.GetTensorShape(fNR);
187 if (fShapeR.size() != 3) {
188 throw std::runtime_error(
"TMVA SOFIE GRU Op input tensor " + fNR +
" is not of 3 dimensions.");
191 if (!model.CheckIfTensorAlreadyExist(fNB)) {
192 throw std::runtime_error(
"TMVA SOFIE GRU op input tensor " + fNB +
" is not found in model.");
194 fShapeB = model.GetTensorShape(fNB);
195 if (fShapeB.size() != 2 && fShapeB.size() != 4) {
196 throw std::runtime_error(
"TMVA SOFIE GRU op input tensor " + fNB +
" is not of 2 or 4 dimensions.");
198 if (fShapeB.size() == 2) {
200 auto original_data = model.GetInitializedTensorData(fNB);
201 size_t num_directions = fShapeW[0];
202 size_t batch_size = (fAttrLayout == 0) ? fShapeX[1] : fShapeX[0];
203 size_t seq_length = (fAttrLayout == 0) ? fShapeX[0] : fShapeX[1];
204 if (fType ==
"float") {
205 float *original_bias =
static_cast<float *
>(original_data.get());
206 float *new_bias =
new float[num_directions * 6 * seq_length * batch_size * fAttrHiddenSize];
207 for (
size_t direction = 0; direction < num_directions; direction++) {
208 for (
size_t i = 0; i < 6; i++) {
209 for (
size_t seq = 0; seq < seq_length; seq++) {
210 for (
size_t batch = 0; batch < batch_size; batch++) {
211 size_t bias_offset = direction * 6 * fAttrHiddenSize + i * fAttrHiddenSize;
212 size_t offset = direction * 6 * batch_size * seq_length * fAttrHiddenSize +
213 i * batch_size * seq_length * fAttrHiddenSize +
214 +seq * batch_size * fAttrHiddenSize + batch * fAttrHiddenSize;
215 std::copy(original_bias + bias_offset, original_bias + bias_offset + fAttrHiddenSize,
222 std::vector<size_t> new_bias_shape = {num_directions, 6, seq_length, batch_size, fAttrHiddenSize};
223 std::shared_ptr<void> new_bias_ptr(new_bias, std::default_delete<
float[]>());
224 model.UpdateInitializedTensor(fNB, model.GetTensorType(fNB), new_bias_shape, new_bias_ptr);
225 fShapeB = model.GetTensorShape(fNB);
229 if (!fNSequence_lens.empty()) {
230 if (!model.CheckIfTensorAlreadyExist(fNSequence_lens)) {
231 throw std::runtime_error(
"TMVA SOFIE GRU Op input tensor " + fNSequence_lens +
"is not found in model.");
233 fShapeSequence_lens = model.GetTensorShape(fNSequence_lens);
234 if (fShapeSequence_lens.size() != 1) {
235 throw std::runtime_error(
"TMVA SOFIE GRU Op input tensor " + fNSequence_lens +
" is not of 1 dimension.");
238 if (!fNInitial_h.empty()) {
239 if (!model.CheckIfTensorAlreadyExist(fNInitial_h)) {
240 throw std::runtime_error(
"TMVA SOFIE GRU Op input tensor " + fNInitial_h +
" is not found in model.");
242 fShapeInitial_h = model.GetTensorShape(fNInitial_h);
243 if (fShapeInitial_h.size() != 3) {
244 throw std::runtime_error(
"TMVA SOFIE GRU Op input tensor " + fNInitial_h +
" is not of 3 dimensions.");
248 fShapeY = ShapeInference({fShapeX, fShapeW})[0];
249 if (!model.CheckIfTensorAlreadyExist(fNY)) {
250 model.AddIntermediateTensor(fNY, model.GetTensorType(fNX), fShapeY);
253 if (!fNY_h.empty()) {
254 fShapeY_h = ShapeInference({fShapeX, fShapeW})[1];
255 if (!model.CheckIfTensorAlreadyExist(fNY_h)) {
256 model.AddIntermediateTensor(fNY_h, model.GetTensorType(fNX), fShapeY_h);
260 for (
auto &activation : fAttrActivations) {
261 if (activation !=
"Relu" && activation !=
"Tanh" && activation !=
"Sigmoid" && activation !=
"Affine" &&
262 activation !=
"LeakyRelu" && activation !=
"ThresholdRelu" && activation !=
"ScaledTanh" &&
263 activation !=
"HardSigmoid" && activation !=
"Elu" && activation !=
"Softsign" && activation !=
"Softplus") {
264 throw std::runtime_error(
"TMVA SOFIE - Activation function " + activation +
" not implemented");
267 if (fAttrDirection ==
"reverse")
268 fAttrDirection =
"backward";
269 if (fAttrDirection !=
"forward" && fAttrDirection !=
"backward" && fAttrDirection !=
"reverse" &&
270 fAttrDirection !=
"bidirectional") {
271 throw std::runtime_error(
"TMVA SOFIE - Invalid GRU direction fAttrDirection = " + fAttrDirection);
273 if (3 * fAttrHiddenSize != fShapeW[1]) {
274 throw std::runtime_error(
"TMVA SOFIE - fAttrHiddenSize must be equal to " + std::to_string(fShapeW[1] / 3));
276 if (fAttrLayout > 1) {
277 throw std::runtime_error(
"TMVA SOFIE - Layout fAttrLayout = " + std::to_string(fAttrLayout) +
278 " must be 0 (timewise) or 1 (batchwise)");
280 if (fAttrLinearBeforeReset > 1) {
281 throw std::runtime_error(
"TMVA SOFIE - fAttrInputForget = " + std::to_string(fAttrLinearBeforeReset) +
284 if (fAttrActivations.empty()) {
285 if (fAttrDirection ==
"bidirectional") {
286 fAttrActivations = {
"Sigmoid",
"Tanh",
"Sigmoid",
"Tanh"};
288 fAttrActivations = {
"Sigmoid",
"Tanh"};
295 std::string opName =
"op_gru_" + fNX;
297 size_t num_directions = fShapeW[0];
298 size_t seq_length = (fAttrLayout == 0) ? fShapeX[0] : fShapeX[1];
299 size_t batch_size = (fAttrLayout == 0) ? fShapeX[1] : fShapeX[0];
300 size_t input_size = fShapeX[2];
302 auto declareVector = [&](std::string
const &
name, std::size_t
n) {
303 std::string fullName = opName +
"_" +
name;
307 if (fAttrLayout != 0) {
308 declareVector(
"input", seq_length * batch_size * input_size);
309 declareVector(
"initial_hidden_state", num_directions * batch_size * fAttrHiddenSize);
310 declareVector(
"initial_cell_state", num_directions * batch_size * fAttrHiddenSize);
313 size_t ff_size = seq_length * batch_size * fAttrHiddenSize;
314 declareVector(
"f_update_gate", ff_size);
315 declareVector(
"f_reset_gate", ff_size);
316 declareVector(
"f_hidden_gate", ff_size);
318 size_t hs_size = seq_length * num_directions * batch_size * fAttrHiddenSize;
319 declareVector(
"update_gate", hs_size);
320 declareVector(
"reset_gate", hs_size);
321 declareVector(
"hidden_gate", hs_size);
324 declareVector(
"feedback", batch_size * fAttrHiddenSize);
327 if (fAttrLayout != 0 || fNY.empty()) {
328 declareVector(
"hidden_state", hs_size);
335 OpName =
"op_" + OpName;
336 std::stringstream out;
338 size_t seq_length = (fAttrLayout == 0) ? fShapeX[0] : fShapeX[1];
339 size_t batch_size = (fAttrLayout == 0) ? fShapeX[1] : fShapeX[0];
340 size_t input_size = fShapeX[2];
341 size_t num_directions = fShapeW[0];
343 auto getVec = [&](std::string
const &
name) {
return "tensor_op_gru_" + fNX +
"_" +
name; };
346 if (fAttrLayout == 0) {
347 out <<
SP << fType <<
" const* " << OpName <<
"_input = tensor_" << fNX <<
";\n";
349 out <<
SP << fType <<
" * " << OpName <<
"_input = " << getVec(
"input") <<
";\n";
350 out <<
SP <<
"for(size_t seq = 0; seq < " << seq_length <<
"; seq++) {\n";
351 out <<
SP <<
SP <<
"for(size_t batch = 0; batch < " << batch_size <<
"; batch++) {\n";
352 out <<
SP <<
SP <<
SP <<
"for(size_t i = 0; i < " << input_size <<
"; i++) {\n";
353 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_input[seq * " << batch_size * input_size <<
" + batch * " << input_size
354 <<
" + i] = " <<
"tensor_" << fNX <<
"[batch * " << seq_length * input_size <<
" + seq * " << input_size
356 out <<
SP <<
SP <<
SP <<
"}\n";
357 out <<
SP <<
SP <<
"}\n";
362 if (!fNInitial_h.empty()) {
363 if (fAttrLayout == 0) {
364 out <<
SP << fType <<
" *" << OpName <<
"_initial_hidden_state = " <<
" tensor_" << fNInitial_h <<
";\n";
366 out <<
SP << fType <<
" * " << OpName <<
"_initial_hidden_state = " << getVec(
"initial_hidden_state") <<
";\n";
367 for (
size_t direction = 0; direction < num_directions; direction++) {
368 out <<
SP <<
"for(size_t batch = 0; batch < " << batch_size <<
"; batch++) {\n";
369 out <<
SP <<
SP <<
"for(size_t h = 0; h < " << fAttrHiddenSize <<
"; h++) {\n";
370 out <<
SP <<
SP <<
SP << OpName <<
"_initial_hidden_state[" << direction * batch_size * fAttrHiddenSize
371 <<
" + batch * " << fAttrHiddenSize <<
" + h] = tensor_" << fNInitial_h <<
"[batch * "
372 << num_directions * fAttrHiddenSize <<
" + " << direction * fAttrHiddenSize <<
" + h];\n";
373 out <<
SP <<
SP <<
"}\n";
380 out <<
SP << fType <<
" * " << OpName <<
"_f_update_gate = " << getVec(
"f_update_gate") <<
";\n";
381 out <<
SP << fType <<
" * " << OpName <<
"_f_reset_gate = " << getVec(
"f_reset_gate") <<
";\n";
382 out <<
SP << fType <<
" * " << OpName <<
"_f_hidden_gate = " << getVec(
"f_hidden_gate") <<
";\n";
384 out <<
SP << fType <<
" * " << OpName <<
"_update_gate = " << getVec(
"update_gate") <<
";\n";
385 out <<
SP << fType <<
" * " << OpName <<
"_reset_gate = " << getVec(
"reset_gate") <<
";\n";
386 out <<
SP << fType <<
" * " << OpName <<
"_hidden_gate = " << getVec(
"hidden_gate") <<
";\n";
388 if (fAttrLayout == 0 && !fNY.empty()) {
389 out <<
SP << fType <<
" *" << OpName <<
"_hidden_state = tensor_" << fNY <<
";\n";
391 out <<
SP << fType <<
" * " << OpName <<
"_hidden_state = " << getVec(
"hidden_state") <<
";\n";
394 out <<
SP << fType <<
" * " << OpName <<
"_feedback = " << getVec(
"feedback") <<
";\n";
396 out <<
SP <<
"char " << OpName <<
"_transA = 'N';\n";
397 out <<
SP <<
"char " << OpName <<
"_transB = 'T';\n";
398 out <<
SP <<
"int " << OpName <<
"_m = " << seq_length * batch_size <<
";\n";
399 out <<
SP <<
"int " << OpName <<
"_m2 = " << batch_size <<
";\n";
400 out <<
SP <<
"int " << OpName <<
"_n = " << fAttrHiddenSize <<
";\n";
401 out <<
SP <<
"int " << OpName <<
"_k = " << input_size <<
";\n";
402 if (fType ==
"float") {
403 out <<
SP <<
"float " << OpName <<
"_alpha = 1.;\n";
404 out <<
SP <<
"float " << OpName <<
"_beta = 0.;\n";
407 out <<
SP <<
"int " << OpName <<
"_bias_size = " << seq_length * batch_size * fAttrHiddenSize <<
";\n";
409 out <<
SP <<
"int " << OpName <<
"_incx = 1;\n";
410 out <<
SP <<
"int " << OpName <<
"_incy = 1;\n";
411 out <<
SP <<
"int " << OpName <<
"_feedback_size = " << batch_size * fAttrHiddenSize <<
";\n";
413 for (
size_t direction = 0; direction < num_directions; direction++) {
414 if (direction == 0) {
415 if (fType ==
"float") {
417 out <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName <<
"_n, &"
418 << OpName <<
"_m, &" << OpName <<
"_k, &" << OpName <<
"_alpha, tensor_" << fNW <<
", &" << OpName
419 <<
"_k, " << OpName <<
"_input, &" << OpName <<
"_k, &" << OpName <<
"_beta, " << OpName
420 <<
"_f_update_gate, &" << OpName <<
"_n);\n";
422 size_t wr_offset = fAttrHiddenSize * input_size;
423 out <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName <<
"_n, &"
424 << OpName <<
"_m, &" << OpName <<
"_k, &" << OpName <<
"_alpha, tensor_" << fNW <<
" + " << wr_offset
425 <<
", &" << OpName <<
"_k, " << OpName <<
"_input, &" << OpName <<
"_k, &" << OpName <<
"_beta, "
426 << OpName <<
"_f_reset_gate, &" << OpName <<
"_n);\n";
428 size_t wh_offset = 2 * fAttrHiddenSize * input_size;
429 out <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName <<
"_n, &"
430 << OpName <<
"_m, &" << OpName <<
"_k, &" << OpName <<
"_alpha, tensor_" << fNW <<
" + " << wh_offset
431 <<
", &" << OpName <<
"_k, " << OpName <<
"_input, &" << OpName <<
"_k, &" << OpName <<
"_beta, "
432 << OpName <<
"_f_hidden_gate, &" << OpName <<
"_n);\n";
435 if (fType ==
"float") {
437 size_t wz_offset = 3 * fAttrHiddenSize * input_size;
438 out <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName <<
"_n, &"
439 << OpName <<
"_m, &" << OpName <<
"_k, &" << OpName <<
"_alpha, tensor_" << fNW <<
" + " << wz_offset
440 <<
", &" << OpName <<
"_k, " << OpName <<
"_input, &" << OpName <<
"_k, &" << OpName <<
"_beta, "
441 << OpName <<
"_f_update_gate, &" << OpName <<
"_n);\n";
443 size_t wr_offset = 3 * fAttrHiddenSize * input_size + fAttrHiddenSize * input_size;
444 out <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName <<
"_n, &"
445 << OpName <<
"_m, &" << OpName <<
"_k, &" << OpName <<
"_alpha, tensor_" << fNW <<
" + " << wr_offset
446 <<
", &" << OpName <<
"_k, " << OpName <<
"_input, &" << OpName <<
"_k, &" << OpName <<
"_beta, "
447 << OpName <<
"_f_reset_gate, &" << OpName <<
"_n);\n";
449 size_t wh_offset = 3 * fAttrHiddenSize * input_size + 2 * fAttrHiddenSize * input_size;
450 out <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName <<
"_n, &"
451 << OpName <<
"_m, &" << OpName <<
"_k, &" << OpName <<
"_alpha, tensor_" << fNW <<
" + " << wh_offset
452 <<
", &" << OpName <<
"_k, " << OpName <<
"_input, &" << OpName <<
"_k, &" << OpName <<
"_beta, "
453 << OpName <<
"_f_hidden_gate, &" << OpName <<
"_n);\n";
458 if (direction == 0) {
459 if (fType ==
"float") {
461 out <<
SP <<
"BLAS::saxpy_(&" << OpName <<
"_bias_size, &" << OpName <<
"_alpha, tensor_" << fNB <<
", &"
462 << OpName <<
"_incx, " << OpName <<
"_f_update_gate, &" << OpName <<
"_incy);\n";
464 size_t rbz_offset = 3 * batch_size * seq_length * fAttrHiddenSize;
465 out <<
SP <<
"BLAS::saxpy_(&" << OpName <<
"_bias_size, &" << OpName <<
"_alpha, tensor_" << fNB <<
" + "
466 << rbz_offset <<
", &" << OpName <<
"_incx, " << OpName <<
"_f_update_gate, &" << OpName
469 size_t wbr_offset = batch_size * seq_length * fAttrHiddenSize;
470 out <<
SP <<
"BLAS::saxpy_(&" << OpName <<
"_bias_size, &" << OpName <<
"_alpha, tensor_" << fNB <<
" + "
471 << wbr_offset <<
", &" << OpName <<
"_incx, " << OpName <<
"_f_reset_gate, &" << OpName
475 size_t rbr_offset = 4 * batch_size * seq_length * fAttrHiddenSize;
476 out <<
SP <<
"BLAS::saxpy_(&" << OpName <<
"_bias_size, &" << OpName <<
"_alpha, tensor_" << fNB <<
" + "
477 << rbr_offset <<
", &" << OpName <<
"_incx, " << OpName <<
"_f_reset_gate, &" << OpName
480 size_t wbh_offset = 2 * batch_size * seq_length * fAttrHiddenSize;
481 out <<
SP <<
"BLAS::saxpy_(&" << OpName <<
"_bias_size, &" << OpName <<
"_alpha, tensor_" << fNB <<
" + "
482 << wbh_offset <<
", &" << OpName <<
"_incx, " << OpName <<
"_f_hidden_gate, &" << OpName
484 if (fAttrLinearBeforeReset == 0) {
486 size_t rbh_offset = 5 * batch_size * seq_length * fAttrHiddenSize;
487 out <<
SP <<
"BLAS::saxpy_(&" << OpName <<
"_bias_size, &" << OpName <<
"_alpha, tensor_" << fNB
488 <<
" + " << rbh_offset <<
", &" << OpName <<
"_incx, " << OpName <<
"_f_hidden_gate, &" << OpName
493 if (fType ==
"float") {
495 size_t wbz_offset = 6 * batch_size * seq_length * fAttrHiddenSize;
496 out <<
SP <<
"BLAS::saxpy_(&" << OpName <<
"_bias_size, &" << OpName <<
"_alpha, tensor_" << fNB <<
" + "
497 << wbz_offset <<
", &" << OpName <<
"_incx, " << OpName <<
"_f_update_gate, &" << OpName
501 size_t rbz_offset = 9 * batch_size * seq_length * fAttrHiddenSize;
502 out <<
SP <<
"BLAS::saxpy_(&" << OpName <<
"_bias_size, &" << OpName <<
"_alpha, tensor_" << fNB <<
" + "
503 << rbz_offset <<
", &" << OpName <<
"_incx, " << OpName <<
"_f_update_gate, &" << OpName
506 size_t wbr_offset = 7 * batch_size * seq_length * fAttrHiddenSize;
507 out <<
SP <<
"BLAS::saxpy_(&" << OpName <<
"_bias_size, &" << OpName <<
"_alpha, tensor_" << fNB <<
" + "
508 << wbr_offset <<
", &" << OpName <<
"_incx, " << OpName <<
"_f_reset_gate, &" << OpName
511 size_t rbr_offset = 10 * batch_size * seq_length * fAttrHiddenSize;
512 out <<
SP <<
"BLAS::saxpy_(&" << OpName <<
"_bias_size, &" << OpName <<
"_alpha, tensor_" << fNB <<
" + "
513 << rbr_offset <<
", &" << OpName <<
"_incx, " << OpName <<
"_f_reset_gate, &" << OpName
516 size_t wbh_offset = 8 * batch_size * seq_length * fAttrHiddenSize;
517 out <<
SP <<
"BLAS::saxpy_(&" << OpName <<
"_bias_size, &" << OpName <<
"_alpha, tensor_" << fNB <<
" + "
518 << wbh_offset <<
", &" << OpName <<
"_incx, " << OpName <<
"_f_hidden_gate, &" << OpName
520 if (fAttrLinearBeforeReset == 0) {
522 size_t rbh_offset = 11 * batch_size * seq_length * fAttrHiddenSize;
523 out <<
SP <<
"BLAS::saxpy_(&" << OpName <<
"_bias_size, &" << OpName <<
"_alpha, tensor_" << fNB
524 <<
" + " << rbh_offset <<
", &" << OpName <<
"_incx, " << OpName <<
"_f_hidden_gate, &" << OpName
532 out <<
SP <<
"for (size_t seq = 0; seq < " << seq_length <<
"; seq++) {\n";
533 out <<
SP <<
SP <<
"size_t offset = seq * " << batch_size * fAttrHiddenSize <<
";\n";
534 if (direction == 0) {
535 out <<
SP <<
SP <<
"size_t gate_offset = seq * " << num_directions * batch_size * fAttrHiddenSize <<
";\n";
537 out <<
SP <<
SP <<
"size_t gate_offset = seq * " << num_directions * batch_size * fAttrHiddenSize <<
" + "
538 << batch_size * fAttrHiddenSize <<
";\n";
540 size_t f_seq_size = batch_size * fAttrHiddenSize;
541 out <<
SP <<
SP <<
"std::copy(" << OpName <<
"_f_update_gate + offset, " << OpName <<
"_f_update_gate + offset + "
542 << f_seq_size <<
", " << OpName <<
"_update_gate + gate_offset);\n";
543 out <<
SP <<
SP <<
"std::copy(" << OpName <<
"_f_reset_gate + offset, " << OpName <<
"_f_reset_gate + offset + "
544 << f_seq_size <<
", " << OpName <<
"_reset_gate + gate_offset);\n";
545 out <<
SP <<
SP <<
"std::copy(" << OpName <<
"_f_hidden_gate + offset, " << OpName <<
"_f_hidden_gate + offset + "
546 << f_seq_size <<
", " << OpName <<
"_hidden_gate + gate_offset);\n";
549 out <<
SP <<
"for (size_t seq = 0; seq < " << seq_length <<
"; seq++) {\n";
550 if (fAttrDirection ==
"backward" || direction == 1) {
551 out <<
SP <<
SP <<
"size_t index = " << seq_length - 1 <<
" - seq;\n";
553 out <<
SP <<
SP <<
"size_t index = seq;\n";
555 out <<
SP <<
SP <<
"int m2 = " << batch_size <<
";\n";
556 if (direction == 0) {
557 out <<
SP <<
SP <<
"size_t offset = index * " << num_directions * batch_size * fAttrHiddenSize <<
";\n";
559 out <<
SP <<
SP <<
"size_t offset = index * " << num_directions * batch_size * fAttrHiddenSize <<
" + "
560 << batch_size * fAttrHiddenSize <<
";\n";
562 size_t size = batch_size * fAttrHiddenSize;
564 out <<
SP <<
SP <<
"if (seq == 0) {\n";
565 if (!fNInitial_h.empty()) {
566 if (direction == 0) {
567 if (fType ==
"float") {
568 out <<
SP <<
SP <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName
569 <<
"_n, &m2, &" << OpName <<
"_n, &" << OpName <<
"_alpha, tensor_" << fNR <<
", &" << OpName
570 <<
"_n, " << OpName <<
"_initial_hidden_state, &" << OpName <<
"_n, &" << OpName <<
"_alpha, "
571 << OpName <<
"_update_gate + offset, &" << OpName <<
"_n);\n";
572 size_t rr_offset = fAttrHiddenSize * fAttrHiddenSize;
573 out <<
SP <<
SP <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName
574 <<
"_n, &m2, &" << OpName <<
"_n, &" << OpName <<
"_alpha, tensor_" << fNR <<
" + " << rr_offset
575 <<
", &" << OpName <<
"_n, " << OpName <<
"_initial_hidden_state, &" << OpName <<
"_n, &" << OpName
576 <<
"_alpha, " << OpName <<
"_reset_gate + offset, &" << OpName <<
"_n);\n";
579 if (fType ==
"float") {
580 size_t rz_offset = 3 * fAttrHiddenSize * fAttrHiddenSize;
581 out <<
SP <<
SP <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName
582 <<
"_n, &m2, &" << OpName <<
"_n, &" << OpName <<
"_alpha, tensor_" << fNR <<
" + " << rz_offset
583 <<
", &" << OpName <<
"_n, " << OpName <<
"_initial_hidden_state, &" << OpName <<
"_n, &" << OpName
584 <<
"_alpha, " << OpName <<
"_update_gate + offset, &" << OpName <<
"_n);\n";
585 size_t rr_offset = 4 * fAttrHiddenSize * fAttrHiddenSize;
586 out <<
SP <<
SP <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName
587 <<
"_n, &m2, &" << OpName <<
"_n, &" << OpName <<
"_alpha, tensor_" << fNR <<
" + " << rr_offset
588 <<
", &" << OpName <<
"_n, " << OpName <<
"_initial_hidden_state, &" << OpName <<
"_n, &" << OpName
589 <<
"_alpha, " << OpName <<
"_reset_gate + offset, &" << OpName <<
"_n);\n";
593 out <<
SP <<
SP <<
"} else {\n";
595 if (direction == 0) {
596 if (fAttrDirection ==
"backward") {
597 out <<
SP <<
SP <<
SP <<
"size_t previous_offset = (index + 1) * "
598 << num_directions * batch_size * fAttrHiddenSize <<
";\n";
600 out <<
SP <<
SP <<
SP <<
"size_t previous_offset = (seq - 1) * "
601 << num_directions * batch_size * fAttrHiddenSize <<
";\n";
603 if (fType ==
"float") {
604 out <<
SP <<
SP <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName
605 <<
"_n, &m2, &" << OpName <<
"_n, &" << OpName <<
"_alpha, tensor_" << fNR <<
", &" << OpName <<
"_n, "
606 << OpName <<
"_hidden_state + previous_offset, &" << OpName <<
"_n, &" << OpName <<
"_alpha, " << OpName
607 <<
"_update_gate + offset, &" << OpName <<
"_n);\n";
608 size_t rr_offset = fAttrHiddenSize * fAttrHiddenSize;
609 out <<
SP <<
SP <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName
610 <<
"_n, &m2, &" << OpName <<
"_n, &" << OpName <<
"_alpha, tensor_" << fNR <<
" + " << rr_offset
611 <<
", &" << OpName <<
"_n, " << OpName <<
"_hidden_state + previous_offset, &" << OpName <<
"_n, &"
612 << OpName <<
"_alpha, " << OpName <<
"_reset_gate + offset, &" << OpName <<
"_n);\n";
615 out <<
SP <<
SP <<
SP <<
"size_t previous_offset = (index + 1) * "
616 << num_directions * batch_size * fAttrHiddenSize <<
" + " << batch_size * fAttrHiddenSize <<
";\n";
617 if (fType ==
"float") {
618 size_t rz_offset = 3 * fAttrHiddenSize * fAttrHiddenSize;
619 out <<
SP <<
SP <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName
620 <<
"_n, &m2, &" << OpName <<
"_n, &" << OpName <<
"_alpha, tensor_" << fNR <<
" + " << rz_offset
621 <<
", &" << OpName <<
"_n, " << OpName <<
"_hidden_state + previous_offset, &" << OpName <<
"_n, &"
622 << OpName <<
"_alpha, " << OpName <<
"_update_gate + offset, &" << OpName <<
"_n);\n";
623 size_t rr_offset = 4 * fAttrHiddenSize * fAttrHiddenSize;
624 out <<
SP <<
SP <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName
625 <<
"_n, &m2, &" << OpName <<
"_n, &" << OpName <<
"_alpha, tensor_" << fNR <<
" + " << rr_offset
626 <<
", &" << OpName <<
"_n, " << OpName <<
"_hidden_state + previous_offset, &" << OpName <<
"_n, &"
627 << OpName <<
"_alpha, " << OpName <<
"_reset_gate + offset, &" << OpName <<
"_n);\n";
630 out <<
SP <<
SP <<
"}\n";
633 if (fAttrClip > .0) {
634 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
635 if (fType ==
"float") {
636 out <<
SP <<
SP <<
SP <<
"float z = (" << OpName <<
"_update_gate[i] > " << -fAttrClip <<
") ? " << OpName
637 <<
"_update_gate[i] : " << -fAttrClip <<
";\n";
639 out <<
SP <<
SP <<
SP << OpName <<
"_update_gate[i] = (z < " << fAttrClip <<
") ? z : " << fAttrClip <<
";\n";
640 if (fType ==
"float") {
641 out <<
SP <<
SP <<
SP <<
"float r = (" << OpName <<
"_reset_gate[i] > " << -fAttrClip <<
") ? " << OpName
642 <<
"_reset_gate[i] : " << -fAttrClip <<
";\n";
644 out <<
SP <<
SP <<
SP << OpName <<
"_reset_gate[i] = (r < " << fAttrClip <<
") ? r : " << fAttrClip <<
";\n";
645 out <<
SP <<
SP <<
"}\n";
649 if (fAttrActivations[direction * 2] ==
"Relu") {
650 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
651 out <<
SP <<
SP <<
SP <<
"if (" << OpName <<
"_update_gate[i] < 0.)\n";
652 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_update_gate[i] = 0.;\n";
653 out <<
SP <<
SP <<
SP <<
"if (" << OpName <<
"_reset_gate[i] < 0.)\n";
654 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_reset_gate[i] = 0.;\n";
655 out <<
SP <<
SP <<
"}\n";
656 }
else if (fAttrActivations[direction * 2] ==
"Tanh") {
657 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
658 if (fType ==
"float") {
659 out <<
SP <<
SP <<
SP <<
"float z = exp(-2 * " << OpName <<
"_update_gate[i]);\n";
661 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_update_gate[i] = (1. - z) / (1. + z);\n";
662 if (fType ==
"float") {
663 out <<
SP <<
SP <<
SP <<
"float r = exp(-2 * " << OpName <<
"_reset_gate[i]);\n";
665 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_reset_gate[i] = (1. - r) / (1. + r);\n";
666 out <<
SP <<
SP <<
"}\n";
667 }
else if (fAttrActivations[direction * 2] ==
"Sigmoid") {
668 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
669 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_update_gate[i] = 1. / (1. + exp(-" << OpName
670 <<
"_update_gate[i]));\n";
671 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_reset_gate[i] = 1. / (1. + exp(-" << OpName
672 <<
"_reset_gate[i]));\n";
673 out <<
SP <<
SP <<
"}\n";
674 }
else if (fAttrActivations[direction * 2] ==
"Affine") {
675 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
676 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_update_gate[i] = " << fAttrActivationAlpha[direction * 2] <<
" * "
677 << OpName <<
"_update_gate[i] + " << fAttrActivationBeta[direction * 2] <<
";\n";
678 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_reset_gate[i] = " << fAttrActivationAlpha[direction * 2] <<
" * "
679 << OpName <<
"_reset_gate[i] + " << fAttrActivationBeta[direction * 2] <<
";\n";
680 out <<
SP <<
SP <<
"}\n";
681 }
else if (fAttrActivations[direction * 2] ==
"ScaledTanh") {
682 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
683 if (fType ==
"float") {
684 out <<
SP <<
SP <<
SP <<
"float z = exp(-2 * " << fAttrActivationBeta[direction * 2] <<
" * " << OpName
685 <<
"_update_gate[i]);\n";
687 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_update_gate[i] = " << fAttrActivationAlpha[direction * 2]
688 <<
" * (1. - z) / (1. + z);\n";
689 if (fType ==
"float") {
690 out <<
SP <<
SP <<
SP <<
"float r = exp(-2 * " << fAttrActivationBeta[direction * 2] <<
" * " << OpName
691 <<
"_reset_gate[i]);\n";
693 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_reset_gate[i] = " << fAttrActivationAlpha[direction * 2]
694 <<
" * (1. - r) / (1. + r);\n";
695 out <<
SP <<
SP <<
"}\n";
696 }
else if (fAttrActivations[direction * 2] ==
"HardSigmoid") {
697 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
698 if (fType ==
"float") {
699 out <<
SP <<
SP <<
SP <<
"float za = " << fAttrActivationAlpha[direction * 2] <<
" * " << OpName
700 <<
"_update_gate[i] + " << fAttrActivationBeta[direction * 2] <<
";\n";
701 out <<
SP <<
SP <<
SP <<
"float zb = (za > 0.) ? za : 0.;\n";
703 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_update_gate[i] = (zb < 1.) ? zb : 1.;\n";
704 if (fType ==
"float") {
705 out <<
SP <<
SP <<
SP <<
"float ra = " << fAttrActivationAlpha[direction * 2] <<
" * " << OpName
706 <<
"_reset_gate[i] + " << fAttrActivationBeta[direction * 2] <<
";\n";
707 out <<
SP <<
SP <<
SP <<
"float rb = (ra > 0.) ? ra : 0.;\n";
709 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_reset_gate[i] = (rb < 1.) ? rb : 1.;\n";
710 out <<
SP <<
SP <<
"}\n";
711 }
else if (fAttrActivations[direction * 2] ==
"LeakyRelu") {
712 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
713 out <<
SP <<
SP <<
SP <<
"if (" << OpName <<
"_update_gate[i] < 0.)\n";
714 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_update_gate[i] = " << fAttrActivationAlpha[direction * 2] <<
" * "
715 << OpName <<
"_update_gate[i];\n";
716 out <<
SP <<
SP <<
SP <<
"if (" << OpName <<
"_reset_gate[i] < 0.)\n";
717 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_reset_gate[i] = " << fAttrActivationAlpha[direction * 2] <<
" * "
718 << OpName <<
"_reset_gate[i];\n";
719 out <<
SP <<
SP <<
"}\n";
720 }
else if (fAttrActivations[direction * 2] ==
"ThresholdRelu") {
721 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
722 out <<
SP <<
SP <<
SP <<
"if (" << OpName <<
"_update_gate[i] < " << fAttrActivationAlpha[direction * 2]
724 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_update_gate[i] = 0.;\n";
725 out <<
SP <<
SP <<
SP <<
"if (" << OpName <<
"_reset_gate[i] < " << fAttrActivationAlpha[direction * 2]
727 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_reset_gate[i] = 0.;\n";
728 out <<
SP <<
SP <<
"}";
729 }
else if (fAttrActivations[direction * 2] ==
"Elu") {
730 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
731 out <<
SP <<
SP <<
SP <<
"if (" << OpName <<
"_update_gate[i] < 0.)\n";
732 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_update_gate[i] = " << fAttrActivationAlpha[direction * 2]
733 <<
" * exp(" << OpName <<
"_update_gate[i] - 1.);\n";
734 out <<
SP <<
SP <<
SP <<
"if (" << OpName <<
"_reset_gate[i] < 0.)\n";
735 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_reset_gate[i] = " << fAttrActivationAlpha[direction * 2]
736 <<
" * exp(" << OpName <<
"_reset_gate[i] - 1.);\n";
737 out <<
SP <<
SP <<
"}\n";
738 }
else if (fAttrActivations[direction * 2] ==
"Softsign") {
739 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
740 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_update_gate[i] = " << OpName <<
"_update_gate[i] / (1. + abs("
741 << OpName <<
"_update_gate[i]));\n";
742 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_reset_gate[i] = " << OpName <<
"_reset_gate[i] / (1. + abs("
743 << OpName <<
"_reset_gate[i]));\n";
744 out <<
SP <<
SP <<
"}\n";
746 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
747 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_update_gate[i] = log(1. + exp(" << OpName <<
"_update_gate[i]));\n";
748 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_reset_gate[i] = log(1. + exp(" << OpName <<
"_reset_gate[i]));\n";
749 out <<
SP <<
SP <<
"}\n";
752 if (fAttrLinearBeforeReset == 0) {
753 out <<
SP <<
SP <<
"if (seq == 0) {\n";
754 if (!fNInitial_h.empty()) {
756 out <<
SP <<
SP <<
SP <<
"for (size_t i = 0; i < " <<
size <<
"; i++) {\n";
757 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_feedback[i] = " << OpName <<
"_reset_gate[i + offset] * "
758 << OpName <<
"_initial_hidden_state[i];\n";
759 out <<
SP <<
SP <<
SP <<
"}\n";
761 out <<
SP <<
SP <<
"} else {\n";
763 if (direction == 0) {
764 if (fAttrDirection ==
"backward") {
765 out <<
SP <<
SP <<
SP <<
"size_t previous_offset = (index + 1) * "
766 << num_directions * batch_size * fAttrHiddenSize <<
";\n";
768 out <<
SP <<
SP <<
SP <<
"size_t previous_offset = (seq - 1) * "
769 << num_directions * batch_size * fAttrHiddenSize <<
";\n";
772 out <<
SP <<
SP <<
SP <<
"size_t previous_offset = (index + 1) * "
773 << num_directions * batch_size * fAttrHiddenSize <<
" + " << batch_size * fAttrHiddenSize <<
";\n";
775 out <<
SP <<
SP <<
SP <<
"for (size_t i = 0; i < " <<
size <<
"; i++) {\n";
776 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_feedback[i] = " << OpName <<
"_reset_gate[i + offset] * " << OpName
777 <<
"_hidden_state[i + previous_offset];\n";
778 out <<
SP <<
SP <<
SP <<
"}\n";
779 out <<
SP <<
SP <<
"}\n";
781 size_t rh_offset = (direction == 0)
782 ? 2 * fAttrHiddenSize * fAttrHiddenSize
783 : 3 * fAttrHiddenSize * fAttrHiddenSize + 2 * fAttrHiddenSize * fAttrHiddenSize;
784 out <<
SP <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName <<
"_n, &"
785 << OpName <<
"_m2, &" << OpName <<
"_n, &" << OpName <<
"_alpha, tensor_" << fNR <<
" + " << rh_offset
786 <<
", &" << OpName <<
"_n, " << OpName <<
"_feedback, &" << OpName <<
"_n, &" << OpName <<
"_beta, "
787 << OpName <<
"_feedback, &" << OpName <<
"_n);\n";
791 size_t rh_offset = (direction == 0)
792 ? 2 * fAttrHiddenSize * fAttrHiddenSize
793 : 3 * fAttrHiddenSize * fAttrHiddenSize + 2 * fAttrHiddenSize * fAttrHiddenSize;
794 out <<
SP <<
SP <<
"if (seq == 0) {\n";
795 if (!fNInitial_h.empty()) {
797 out <<
SP <<
SP <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName
798 <<
"_n, &" << OpName <<
"_m2, &" << OpName <<
"_n, &" << OpName <<
"_alpha, tensor_" << fNR <<
" + "
799 << rh_offset <<
", &" << OpName <<
"_n, " << OpName <<
"_initial_hidden_state, &" << OpName <<
"_n, &"
800 << OpName <<
"_beta, " << OpName <<
"_feedback, &" << OpName <<
"_n);\n";
802 out <<
SP <<
SP <<
"} else {\n";
804 if (direction == 0) {
805 if (fAttrDirection ==
"backward") {
806 out <<
SP <<
SP <<
SP <<
"size_t previous_offset = (index + 1) * "
807 << num_directions * batch_size * fAttrHiddenSize <<
";\n";
809 out <<
SP <<
SP <<
SP <<
"size_t previous_offset = (seq - 1) * "
810 << num_directions * batch_size * fAttrHiddenSize <<
";\n";
813 out <<
SP <<
SP <<
SP <<
"size_t previous_offset = (index + 1) * "
814 << num_directions * batch_size * fAttrHiddenSize <<
" + " << batch_size * fAttrHiddenSize <<
";\n";
816 out <<
SP <<
SP <<
SP <<
"BLAS::sgemm_(&" << OpName <<
"_transB, &" << OpName <<
"_transA, &" << OpName
817 <<
"_n, &" << OpName <<
"_m2, &" << OpName <<
"_n, &" << OpName <<
"_alpha, tensor_" << fNR <<
" + "
818 << rh_offset <<
", &" << OpName <<
"_n, " << OpName <<
"_hidden_state + previous_offset, &" << OpName
819 <<
"_n, &" << OpName <<
"_beta, " << OpName <<
"_feedback, &" << OpName <<
"_n);\n";
821 out <<
SP <<
SP <<
"}\n";
824 size_t rbh_offset = (direction == 0) ? 5 * batch_size * seq_length * fAttrHiddenSize
825 : 11 * batch_size * seq_length * fAttrHiddenSize;
826 out <<
SP <<
SP <<
"BLAS::saxpy_(&" << OpName <<
"_feedback_size, &" << OpName <<
"_alpha, tensor_" << fNB
827 <<
" + " << rbh_offset <<
", &" << OpName <<
"_incx, " << OpName <<
"_feedback, &" << OpName
831 out <<
SP <<
SP <<
"for (size_t i = 0; i < " <<
size <<
"; i++) {\n";
832 out <<
SP <<
SP <<
SP << OpName <<
"_feedback[i] *= " << OpName <<
"_reset_gate[i + offset];\n";
833 out <<
SP <<
SP <<
"}\n";
837 out <<
SP <<
SP <<
"BLAS::saxpy_(&" << OpName <<
"_feedback_size, &" << OpName <<
"_alpha, " << OpName
838 <<
"_feedback, &" << OpName <<
"_incx, " << OpName <<
"_hidden_gate + offset, &" << OpName <<
"_incy);\n";
841 if (fAttrClip > .0) {
842 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
843 if (fType ==
"float") {
844 out <<
SP <<
SP <<
SP <<
"float x = (" << OpName <<
"_hidden_gate[i] > " << -fAttrClip <<
") ? " << OpName
845 <<
"_hidden_gate[i] : " << -fAttrClip <<
";\n";
847 out <<
SP <<
SP <<
SP << OpName <<
"_hidden_gate[i] = (x < " << fAttrClip <<
") ? x : " << fAttrClip <<
";\n";
848 out <<
SP <<
SP <<
"}\n";
852 if (fAttrActivations[direction * 2 + 1] ==
"Relu") {
853 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
854 out <<
SP <<
SP <<
SP <<
"if (" << OpName <<
"_hidden_gate[i] < 0.)\n";
855 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_hidden_gate[i] = 0.;\n";
856 out <<
SP <<
SP <<
"}\n";
857 }
else if (fAttrActivations[direction * 2 + 1] ==
"Tanh") {
858 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
859 if (fType ==
"float") {
860 out <<
SP <<
SP <<
SP <<
"float ex = exp(-2 * " << OpName <<
"_hidden_gate[i]);\n";
862 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_hidden_gate[i] = (1. - ex) / (1. + ex);\n";
863 out <<
SP <<
SP <<
"}\n";
864 }
else if (fAttrActivations[direction * 2 + 1] ==
"Sigmoid") {
865 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
866 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_hidden_gate[i] = 1. / (1. + exp(-" << OpName
867 <<
"_hidden_gate[i]));\n";
868 out <<
SP <<
SP <<
"}\n";
869 }
else if (fAttrActivations[direction * 2 + 1] ==
"Affine") {
870 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
871 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_hidden_gate[i] = " << fAttrActivationAlpha[direction * 2 + 1]
872 <<
" * " << OpName <<
"_hidden_gate[i] + " << fAttrActivationBeta[direction * 2 + 1] <<
";\n";
873 out <<
SP <<
SP <<
"}\n";
874 }
else if (fAttrActivations[direction * 2 + 1] ==
"ScaledTanh") {
875 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
876 if (fType ==
"float") {
877 out <<
SP <<
SP <<
SP <<
"float ex = exp(-2 * " << fAttrActivationBeta[direction * 2 + 1] <<
" * " << OpName
878 <<
"_hidden_gate[i]);\n";
880 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_hidden_gate[i] = " << fAttrActivationAlpha[direction * 2 + 1]
881 <<
" * (1. - ex) / (1. + ex);\n";
882 out <<
SP <<
SP <<
"}\n";
883 }
else if (fAttrActivations[direction * 2 + 1] ==
"HardSigmoid") {
884 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
885 if (fType ==
"float") {
886 out <<
SP <<
SP <<
SP <<
"float a = " << fAttrActivationAlpha[direction * 2 + 1] <<
" * " << OpName
887 <<
"_hidden_gate[i] + " << fAttrActivationBeta[direction * 2 + 1] <<
";\n";
888 out <<
SP <<
SP <<
SP <<
"float b = (a > 0.) ? a : 0.;\n";
890 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_hidden_gate[i] = (b < 1.) ? b : 1.;\n";
891 out <<
SP <<
SP <<
"}\n";
892 }
else if (fAttrActivations[direction * 2 + 1] ==
"LeakyRelu") {
893 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
894 out <<
SP <<
SP <<
SP <<
"if (" << OpName <<
"_hidden_gate[i] < 0.)\n";
895 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_hidden_gate[i] = " << fAttrActivationAlpha[direction * 2 + 1]
896 <<
" * " << OpName <<
"_hidden_gate[i];\n";
897 out <<
SP <<
SP <<
"}\n";
898 }
else if (fAttrActivations[direction * 2 + 1] ==
"ThresholdRelu") {
899 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
900 out <<
SP <<
SP <<
SP <<
"if (" << OpName <<
"_hidden_gate[i] < " << fAttrActivationAlpha[direction * 2 + 1]
902 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_hidden_gate[i] = 0.;\n";
903 out <<
SP <<
SP <<
"}";
904 }
else if (fAttrActivations[direction * 2 + 1] ==
"Elu") {
905 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
906 out <<
SP <<
SP <<
SP <<
"if (" << OpName <<
"_hidden_gate[i] < 0.)\n";
907 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_hidden_gate[i] = " << fAttrActivationAlpha[direction * 2 + 1]
908 <<
" * exp(" << OpName <<
"_hidden_gate[i] - 1.);\n";
909 out <<
SP <<
SP <<
"}\n";
910 }
else if (fAttrActivations[direction * 2 + 1] ==
"Softsign") {
911 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
912 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_hidden_gate[i] = " << OpName <<
"_hidden_gate[i] / (1. + abs("
913 << OpName <<
"_hidden_gate[i]));\n";
914 out <<
SP <<
SP <<
"}\n";
916 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
917 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_hidden_gate[i] = log(1. + exp(" << OpName <<
"_hidden_gate[i]));\n";
918 out <<
SP <<
SP <<
"}\n";
922 out <<
SP <<
SP <<
"for (size_t i = offset; i < offset + " <<
size <<
"; i++) {\n";
923 out <<
SP <<
SP <<
SP << OpName <<
"_hidden_state[i] = ( 1. - " << OpName <<
"_update_gate[i]) * " << OpName
924 <<
"_hidden_gate[i];\n";
925 out <<
SP <<
SP <<
"}\n";
927 out <<
SP <<
SP <<
"if (seq == 0) {\n";
928 if (!fNInitial_h.empty()) {
930 out <<
SP <<
SP <<
SP <<
"for (size_t i = 0; i < " <<
size <<
"; i++) {\n";
931 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_hidden_state[i + offset] += " << OpName
932 <<
"_update_gate[i + offset] * " << OpName <<
"_initial_hidden_state[i];\n";
933 out <<
SP <<
SP <<
SP <<
"}\n";
935 out <<
SP <<
SP <<
"} else {\n";
937 if (direction == 0) {
938 if (fAttrDirection ==
"backward") {
939 out <<
SP <<
SP <<
SP <<
"size_t previous_offset = (index + 1) * "
940 << num_directions * batch_size * fAttrHiddenSize <<
";\n";
942 out <<
SP <<
SP <<
SP <<
"size_t previous_offset = (seq - 1) * "
943 << num_directions * batch_size * fAttrHiddenSize <<
";\n";
946 out <<
SP <<
SP <<
SP <<
"size_t previous_offset = (index + 1) * "
947 << num_directions * batch_size * fAttrHiddenSize <<
" + " << batch_size * fAttrHiddenSize <<
";\n";
949 out <<
SP <<
SP <<
SP <<
"for (size_t i = 0; i < " <<
size <<
"; i++) {\n";
950 out <<
SP <<
SP <<
SP <<
SP << OpName <<
"_hidden_state[i + offset] += " << OpName
951 <<
"_update_gate[i + offset] * " << OpName <<
"_hidden_state[i + previous_offset];\n";
952 out <<
SP <<
SP <<
SP <<
"}\n";
953 out <<
SP <<
SP <<
"}\n";
959 if (!fNSequence_lens.empty()) {
960 out <<
SP <<
"for (size_t seq = 0; seq < " << seq_length <<
"; seq++) {\n";
961 out <<
SP <<
SP <<
"for (size_t batch = 0; batch < " << batch_size <<
"; batch++) {\n";
962 out <<
SP <<
SP <<
SP <<
"if (seq >= tensor_" << fNSequence_lens <<
"[batch]) {\n";
963 for (
size_t direction = 0; direction < num_directions; direction++) {
964 out <<
SP <<
SP <<
SP <<
SP <<
SP <<
"for (size_t h = 0; h < " << fAttrHiddenSize <<
"; h++) {\n";
965 out <<
SP <<
SP <<
SP <<
SP <<
SP <<
SP << OpName <<
"_hidden_state[seq * "
966 << num_directions * batch_size * fAttrHiddenSize + direction * batch_size * fAttrHiddenSize
967 <<
" + batch * " << fAttrHiddenSize <<
" + h] = 0.;\n";
970 out <<
SP <<
SP <<
SP <<
"}\n";
971 out <<
SP <<
SP <<
"}\n";
976 if (fAttrLayout == 0) {
977 if (!fNY_h.empty()) {
979 if (fNSequence_lens.empty()) {
980 size_t yh_size = batch_size * fAttrHiddenSize;
981 if (fAttrDirection ==
"backward") {
982 out <<
SP <<
"std::copy(" << OpName <<
"_hidden_state, " << OpName <<
"_hidden_state + " << yh_size
983 <<
", tensor_" << fNY_h <<
");\n";
985 size_t offset = (seq_length - 1) * num_directions * batch_size * fAttrHiddenSize;
986 out <<
SP <<
"std::copy(" << OpName <<
"_hidden_state + " <<
offset <<
", " << OpName
987 <<
"_hidden_state + " <<
offset <<
" + " << yh_size <<
", tensor_" << fNY_h <<
");\n";
989 if (num_directions == 2) {
990 out <<
SP <<
"std::copy(" << OpName <<
"_hidden_state + " << yh_size <<
", " << OpName
991 <<
"_hidden_state + " << 2 * yh_size <<
", tensor_" << fNY_h <<
" + " << yh_size <<
");\n";
994 if (fAttrDirection ==
"backward") {
995 out <<
SP <<
"for (size_t batch = 0; batch < " << batch_size <<
"; batch++) {\n";
996 out <<
SP <<
SP <<
"size_t offset = batch * " << fAttrHiddenSize <<
";\n";
997 out <<
SP <<
SP <<
"std::copy(" << OpName <<
"_hidden_state + offset, " << OpName
998 <<
"_hidden_state + offset + " << fAttrHiddenSize <<
", tensor_" << fNY_h <<
" + offset);\n";
1001 out <<
SP <<
"for (size_t batch = 0; batch < " << batch_size <<
"; batch++) {\n";
1002 out <<
SP <<
SP <<
"size_t seq = " <<
"tensor_" << fNSequence_lens <<
"[batch] - 1;\n";
1003 out <<
SP <<
SP <<
"size_t offset = seq * " << num_directions * batch_size * fAttrHiddenSize
1004 <<
" + batch * " << fAttrHiddenSize <<
";\n";
1005 out <<
SP <<
SP <<
"size_t yh_offset = batch * " << fAttrHiddenSize <<
";\n";
1006 out <<
SP <<
SP <<
"std::copy(" << OpName <<
"_hidden_state + offset, " << OpName
1007 <<
"_hidden_state + offset + " << fAttrHiddenSize <<
", tensor_" << fNY_h <<
" + yh_offset);\n";
1010 if (num_directions == 2) {
1011 out <<
SP <<
"for (size_t batch = 0; batch < " << batch_size <<
"; batch++) {\n";
1012 out <<
SP <<
SP <<
"size_t offset = " << batch_size * fAttrHiddenSize <<
" + batch * " << fAttrHiddenSize
1014 out <<
SP <<
SP <<
"size_t yh_offset = " << batch_size * fAttrHiddenSize <<
" + batch * "
1015 << fAttrHiddenSize <<
";\n";
1016 out <<
SP <<
SP <<
"std::copy(" << OpName <<
"_hidden_state + offset, " << OpName
1017 <<
"_hidden_state + offset + " << fAttrHiddenSize <<
", tensor_" << fNY_h <<
" + yh_offset);\n";
1025 for (
size_t direction = 0; direction < num_directions; direction++) {
1026 out <<
SP <<
"for (size_t seq = 0; seq < " << seq_length <<
"; seq++) {\n";
1027 out <<
SP <<
SP <<
"for (size_t batch = 0; batch < " << batch_size <<
"; batch++) {\n";
1028 out <<
SP <<
SP <<
SP <<
"size_t offset = seq * " << num_directions * batch_size * fAttrHiddenSize <<
" + "
1029 << direction * batch_size * fAttrHiddenSize <<
" + batch * " << fAttrHiddenSize <<
";\n";
1030 out <<
SP <<
SP <<
SP <<
"size_t y_offset = batch * " << seq_length * num_directions * fAttrHiddenSize
1031 <<
" + seq * " << num_directions * fAttrHiddenSize <<
" + " << direction * fAttrHiddenSize <<
";\n";
1032 out <<
SP <<
SP <<
SP <<
"std::copy(" << OpName <<
"_hidden_state + offset, " << OpName
1033 <<
"_hidden_state + offset + " << fAttrHiddenSize <<
", tensor_" << fNY <<
" + y_offset);\n";
1034 out <<
SP <<
SP <<
"}\n";
1038 if (!fNY_h.empty()) {
1040 if (fAttrDirection ==
"backward") {
1041 out <<
SP <<
"for (size_t batch = 0; batch < " << batch_size <<
"; batch++) {\n";
1042 out <<
SP <<
SP <<
"size_t offset = batch * " << fAttrHiddenSize <<
";\n";
1043 out <<
SP <<
SP <<
"size_t yh_offset = batch * " << num_directions * fAttrHiddenSize <<
";\n";
1044 out <<
SP <<
SP <<
"std::copy(" << OpName <<
"_hidden_state + offset, " << OpName
1045 <<
"_hidden_state + offset + " << fAttrHiddenSize <<
", tensor_" << fNY_h <<
" + yh_offset);\n";
1048 out <<
SP <<
"for (size_t batch = 0; batch < " << batch_size <<
"; batch++) {\n";
1049 if (fNSequence_lens.empty()) {
1050 out <<
SP <<
SP <<
"size_t seq = " << seq_length - 1 <<
";\n";
1052 out <<
SP <<
SP <<
"size_t seq = " <<
"tensor_" << fNSequence_lens <<
"[batch] - 1;\n";
1054 out <<
SP <<
SP <<
"size_t offset = seq * " << num_directions * batch_size * fAttrHiddenSize
1055 <<
" + batch * " << fAttrHiddenSize <<
";\n";
1056 out <<
SP <<
SP <<
"size_t yh_offset = batch * " << num_directions * fAttrHiddenSize <<
";\n";
1057 out <<
SP <<
SP <<
"std::copy(" << OpName <<
"_hidden_state + offset, " << OpName
1058 <<
"_hidden_state + offset + " << fAttrHiddenSize <<
", tensor_" << fNY_h <<
" + yh_offset);\n";
1061 if (num_directions == 2) {
1062 out <<
SP <<
"for (size_t batch = 0; batch < " << batch_size <<
"; batch++) {\n";
1063 out <<
SP <<
SP <<
"size_t offset = " << batch_size * fAttrHiddenSize <<
" + batch * " << fAttrHiddenSize
1065 out <<
SP <<
SP <<
"size_t yh_offset = batch * " << num_directions * fAttrHiddenSize <<
" + "
1066 << fAttrHiddenSize <<
";\n";
1067 out <<
SP <<
SP <<
"std::copy(" << OpName <<
"_hidden_state + offset, " << OpName
1068 <<
"_hidden_state + offset + " << fAttrHiddenSize <<
", tensor_" << fNY_h <<
" + yh_offset);\n";