Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
ROperator_LSTM.hxx
Go to the documentation of this file.
1#ifndef TMVA_SOFIE_ROPERATOR_LSTM
2#define TMVA_SOFIE_ROPERATOR_LSTM
3
4#include "TMVA/RModel.hxx"
5#include "TMVA/ROperator.hxx"
7
8#include <memory>
9#include <sstream>
10#include <string>
11#include <vector>
12
14
15/*! \brief Long Short-Term Memory operator
16 *
17 * Inference code generation for one-layer LSTM. Supports forward, reverse and bidirectional LSTM.
18 * See the <a href="https://github.com/onnx/onnx/blob/master/docs/Operators.md#LSTM">ONNX documentation</a>
19 * for details about the supported LSTM architectures.
20 */
21template <typename T> class ROperator_LSTM final : public ROperator {
22 private:
23 std::vector<float> fAttrActivationAlpha; ///< Sacling values used by some activation functions
24 std::vector<float> fAttrActivationBeta; ///< Scaling values used by some activation functions
25 std::vector<std::string> fAttrActivations; ///< Activation functions
26 float fAttrClip; ///< Clip threshold
27 std::string fAttrDirection; ///< Direction of processing
28 size_t fAttrHiddenSize; ///< Number of the hidden layers
29 size_t fAttrInputForget; ///< Forget gate
30 size_t fAttrLayout; ///< Data layout
31
32 std::string fNX; ///< Name of the input
33 std::string fNW; ///< Name of the weights
34 std::string fNR; ///< Name of the recurrence
35 std::string fNB; ///< Name of the bias
36 std::string fNSequence_lens; ///< Name of length of the sequences
37 std::string fNInitial_h; ///< Name of the initial value of the hidden states
38 std::string fNInitial_c; ///< Name of the initial value of the cell states
39 std::string fNP; ///< Name of peepholes
40 std::string fNY; ///< Name of the output
41 std::string fNY_h; ///< Name of the last sequence of the output
42 std::string fNY_c; ///< Name of the last sequence of the cell states
43
44 std::vector<size_t> fShapeX; ///< Shape of the input
45 std::vector<size_t> fShapeW; ///< Shape of the weights
46 std::vector<size_t> fShapeR; ///< Shape of the recurrence
47 std::vector<size_t> fShapeB; ///< Shape of the bias
48 std::vector<size_t> fShapeSequence_lens; ///< Shape of the length of the sequences
49 std::vector<size_t> fShapeInitial_h; ///< Shape of the initial value of the hidden states
50 std::vector<size_t> fShapeInitial_c; ///< Shape of the initial value of the cell states
51 std::vector<size_t> fShapeP; ///< Shape of the peepholes
52 std::vector<size_t> fShapeY; ///< Shape of the output
53 std::vector<size_t> fShapeY_h; ///< Shape of the last sequence of the output
54 std::vector<size_t> fShapeY_c; ///< Shape of the last sequence of the cell states
55
56 std::string fType; ///< Type of the tensors
57
58 public:
59 /*! Default constructor of ROperator_LSTM */
61
62 /*! \brief Constructor of ROperator_LSTM from the attributes
63 *
64 * \param activation_alpha scaling values used by some activation functions
65 * \param activation_beta scaling values used by some activation functions
66 * \param activations activation functions
67 * \param clip clip threshold
68 * \param direction direction of processing of the sequneces
69 * \param hidden_size number of hidden layers
70 * \param input_forget forget gate
71 * \param layout data layout
72 * \param nameX name of the input tensor
73 * \param nameW name of the weight tensor
74 * \param nameR name of the recurrence tensor
75 * \param nameB name of the bias tensor
76 * \param nameSequence_lens name of the length of the sequences
77 * \param nameInitial_h name of the initial value of the hidden states
78 * \param nameInitial_c name of the initial value of the cell states
79 * \param nameP name of the peepholes tensor
80 * \param nameY name of the output
81 * \param nameY_h name of the last sequence of the output
82 * \param nameY_c name of the last sequence of the cell states
83 */
84 ROperator_LSTM(std::vector<float> activation_alpha,
85 std::vector<float> activation_beta,
86 std::vector<std::string> activations, float clip,
87 std::string direction, size_t hidden_size,
88 size_t input_forget, size_t layout,
89 std::string nameX, std::string nameW, std::string nameR,
90 std::string nameB, std::string nameSequence_lens,
91 std::string nameInitial_h, std::string nameInitial_c, std::string nameP,
92 std::string nameY, std::string nameY_h, std::string nameY_c)
97 fNX(UTILITY::Clean_name(nameX)), fNW(UTILITY::Clean_name(nameW)),
98 fNR(UTILITY::Clean_name(nameR)), fNB(UTILITY::Clean_name(nameB)),
99 fNSequence_lens(UTILITY::Clean_name(nameSequence_lens)),
100 fNInitial_h(UTILITY::Clean_name(nameInitial_h)),
101 fNInitial_c(UTILITY::Clean_name(nameInitial_c)), fNP(UTILITY::Clean_name(nameP)),
102 fNY(UTILITY::Clean_name(nameY)), fNY_h(UTILITY::Clean_name(nameY_h)),
103 fNY_c(UTILITY::Clean_name(nameY_c)) {
104 if (std::is_same<T, float>::value) {
105 fType = "float";
106 } else {
107 throw std::runtime_error(
108 "TMVA SOFIE Encountered unsupported type parsing a LSTM operator");
109 }
110
112 if (!fNB.empty()){
113 fInputTensorNames.emplace_back(fNB);
114 }
115 if (!fNSequence_lens.empty()){
117 }
118 if (!fNInitial_h.empty()){
119 fInputTensorNames.emplace_back(fNInitial_h);
120 }
121 if (!fNInitial_c.empty()){
122 fInputTensorNames.emplace_back(fNInitial_c);
123 }
124 if (!fNP.empty()){
125 fInputTensorNames.emplace_back(fNP);
126 }
127
128 fOutputTensorNames = { };
129 if (!fNY.empty()){
130 fOutputTensorNames.emplace_back(fNY);
131 }
132 if (!fNY_h.empty()){
133 fOutputTensorNames.emplace_back(fNY_h);
134 }
135 if (!fNY_c.empty()){
136 fOutputTensorNames.emplace_back(fNY_c);
137 }
138 }
139
140 /*! \brief Infers the type of the output tensors
141 *
142 * \param input type of the input tensors
143 */
144 std::vector<ETensorType> TypeInference(std::vector<ETensorType> input) override;
145
146 /*! \brief Infers the shape of the output tensors
147 *
148 * \param input shape of the input tensors
149 */
150 std::vector<std::vector<size_t>>
151 ShapeInference(std::vector<std::vector<size_t>> input) override;
152
153 /*! \brief Initialize the model
154 *
155 * \param model Model
156 */
157 void Initialize(RModel &) override;
158
159 /*! \brief Generate the inference code
160 *
161 * \param OpName name of the operator
162 */
163 std::string Generate(std::string OpName) override;
164
165 /*! \brief Generate the code for the Session internal data vectors
166 *
167 * \param opName name of the operator
168 */
169 std::string GenerateSessionMembersCode(std::string opName) override;
170
171 /*! \brief Returns the blas routines needed to compile the generated code
172 */
173 std::vector<std::string> GetBlasRoutines() override { return { std::string("Gemm"), std::string("Axpy") }; }
174};
175
176template <typename T>
177auto ROperator_LSTM<T>::TypeInference(std::vector<ETensorType> input) -> std::vector<ETensorType>
178{
179 ETensorType out = input[0];
180 return {out, out};
181}
182
183template <typename T>
184auto ROperator_LSTM<T>::ShapeInference(std::vector<std::vector<size_t>> input) -> std::vector<std::vector<size_t>>
185{
186 size_t num_directions = input[1][0];
187 size_t hidden_size = input[1][1] / 4;
188 if (fAttrLayout == 0) {
189 size_t seq_length = input[0][0];
190 size_t batch_size = input[0][1];
191 std::vector<std::vector<size_t>> ret({{seq_length, num_directions, batch_size, hidden_size},
194 return ret;
195 } else {
196 size_t batch_size = input[0][0];
197 size_t seq_length = input[0][1];
198 std::vector<std::vector<size_t>> ret({{batch_size, seq_length, num_directions, hidden_size},
201 return ret;
202 }
203}
204
205template <typename T>
207{
208 // Check the input and output tensors
209 if (!model.CheckIfTensorAlreadyExist(fNX)) {
210 throw std::runtime_error("TMVA SOFIE LSTM Op input tensor " + fNX + " is not found in model.");
211 }
212 fShapeX = model.GetTensorShape(fNX);
213 if (fShapeX.size() != 3) {
214 throw std::runtime_error("TMVA SOFIE LSTM Op input tensor " + fNX + " is not of 3 dimensions.");
215 }
216 if (!model.CheckIfTensorAlreadyExist(fNW)) {
217 throw std::runtime_error("TMVA SOFIE LSTM Op input tensor " + fNW + " is not found in model.");
218 }
219 fShapeW = model.GetTensorShape(fNW);
220 if (fShapeW.size() != 3) {
221 throw std::runtime_error("TMVA SOFIE LSTM Op input tensor " + fNW + " is not of 3 dimensions.");
222 }
223 if (!model.CheckIfTensorAlreadyExist(fNR)) {
224 throw std::runtime_error("TMVA SOFIE LSTM Op input tensor " + fNR + " is not found in model.");
225 }
226 fShapeR = model.GetTensorShape(fNR);
227 if (fShapeR.size() != 3) {
228 throw std::runtime_error("TMVA SOFIE LSTM Op input tensor " + fNR + " is not of 3 dimensions.");
229 }
230 if (!fNB.empty()) {
231 if (!model.CheckIfTensorAlreadyExist(fNB)) {
232 throw std::runtime_error("TMVA SOFIE LSTM op input tensor " + fNB + " is not found in model.");
233 }
234 fShapeB = model.GetTensorShape(fNB);
235 if (fShapeB.size() != 2 && fShapeB.size() != 5) {
236 throw std::runtime_error("TMVA SOFIE LSTM op input tensor " + fNB + " is not of 2 or 5 dimensions.");
237 }
238 if (fShapeB.size() == 2) {
239 // Broadcasting the bias
240 auto original_data = model.GetInitializedTensorData(fNB);
241 size_t num_directions = fShapeW[0];
242 size_t seq_length = (fAttrLayout == 0) ? fShapeX[0] : fShapeX[1];
243 size_t batch_size = (fAttrLayout == 0) ? fShapeX[1] : fShapeX[0];
244 if (fType == "float") {
245 float *original_bias = static_cast<float *>(original_data.get());
246 float *new_bias = new float[4 * num_directions * seq_length * batch_size * fAttrHiddenSize];
247 for (size_t gate = 0; gate < 4; gate++) {
248 std::vector<float> sum(fAttrHiddenSize);
249 for (size_t direction = 0; direction < num_directions; direction++) {
250 size_t offset = direction * 8 * fAttrHiddenSize + gate * fAttrHiddenSize;
251 for (size_t h = 0; h < fAttrHiddenSize; h++) {
252 sum[h] = original_bias[offset + h] + original_bias[offset + h + 4 * fAttrHiddenSize];
253 }
254 for (size_t seq = 0; seq < seq_length; seq++) {
255 for (size_t batch = 0; batch < batch_size; batch++) {
256 size_t bias_offset = gate * num_directions * seq_length * batch_size * fAttrHiddenSize +
257 direction * seq_length * batch_size * fAttrHiddenSize +
258 seq * batch_size * fAttrHiddenSize + batch * fAttrHiddenSize;
259 std::copy(sum.begin(), sum.end(), new_bias + bias_offset);
260 }
261 }
262 }
263 }
264 std::vector<size_t> new_bias_shape = {4, num_directions, seq_length, batch_size, fAttrHiddenSize};
265 std::shared_ptr<void> new_bias_ptr(new_bias, std::default_delete<float[]>());
266 model.UpdateInitializedTensor(fNB, model.GetTensorType(fNB), new_bias_shape, new_bias_ptr);
267 fShapeB = model.GetTensorShape(fNB);
268 }
269 }
270 }
271 if (!fNSequence_lens.empty()) {
272 if (!model.CheckIfTensorAlreadyExist(fNSequence_lens)) {
273 throw std::runtime_error("TMVA SOFIE LSTM Op input tensor " + fNSequence_lens + "is not found in model.");
274 }
275 fShapeSequence_lens = model.GetTensorShape(fNSequence_lens);
276 if (fShapeSequence_lens.size() != 1) {
277 throw std::runtime_error("TMVA SOFIE LSTM Op input tensor " + fNSequence_lens + " is not of 1 dimension.");
278 }
279 }
280 if (!fNInitial_h.empty()) {
281 if (!model.CheckIfTensorAlreadyExist(fNInitial_h)) {
282 throw std::runtime_error("TMVA SOFIE LSTM Op input tensor " + fNInitial_h + " is not found in model.");
283 }
284 fShapeInitial_h = model.GetTensorShape(fNInitial_h);
285 if (fShapeInitial_h.size() != 3) {
286 throw std::runtime_error("TMVA SOFIE LSTM Op input tensor " + fNInitial_h + " is not of 3 dimensions.");
287 }
288 }
289 if (!fNInitial_c.empty()) {
290 if (!model.CheckIfTensorAlreadyExist(fNInitial_c)) {
291 throw std::runtime_error("TMVA SOFIE LSTM Op input tensor " + fNInitial_c + " is not found in model.");
292 }
293 fShapeInitial_c = model.GetTensorShape(fNInitial_c);
294 if (fShapeInitial_c.size() != 3) {
295 throw std::runtime_error("TMVA SOFIE LSTM Op input tensor " + fNInitial_c + " is not of 3 dimensions.");
296 }
297 }
298 if (!fNP.empty()) {
299 if (!model.CheckIfTensorAlreadyExist(fNP)) {
300 throw std::runtime_error("TMVA SOFIE LSTM op input tensor " + fNP + " is not found in model.");
301 }
302 fShapeP = model.GetTensorShape(fNP);
303 if (fShapeP.size() != 2 && fShapeP.size() != 4) {
304 throw std::runtime_error("TMVA SOFIE LSTM op input tensor " + fNP + " is not of 2 or 4 dimensions.");
305 }
306 if (fShapeP.size() == 2) {
307 // Broadcasting the weight for peepholes
308 auto original_data = model.GetInitializedTensorData(fNP);
309 size_t num_directions = fShapeW[0];
310 size_t batch_size = (fAttrLayout == 0) ? fShapeX[1] : fShapeX[0];
311 if (fType == "float") {
312 float *original_p = static_cast<float *>(original_data.get());
313 float *new_p = new float[num_directions * 3 * batch_size * fAttrHiddenSize];
314 for (size_t direction = 0; direction < num_directions; direction++) {
315 for (size_t gate = 0; gate < 3; gate++) {
316 size_t p_offset = direction * 3 * fAttrHiddenSize + gate * fAttrHiddenSize;
317 for (size_t batch = 0; batch < batch_size; batch++) {
318 size_t offset = direction * 3 * batch_size * fAttrHiddenSize +
319 gate * batch_size * fAttrHiddenSize + batch * fAttrHiddenSize;
320 std::copy(original_p + p_offset, original_p + p_offset + fAttrHiddenSize, new_p + offset);
321 }
322 }
323 }
324 std::vector<size_t> new_p_shape = {num_directions, 3, batch_size, fAttrHiddenSize};
325 std::shared_ptr<void> new_p_ptr(new_p, std::default_delete<float[]>());
326 model.UpdateInitializedTensor(fNP, model.GetTensorType(fNP), new_p_shape, new_p_ptr);
327 fShapeP = model.GetTensorShape(fNP);
328 }
329 }
330 }
331 if (!fNY.empty()) {
332 fShapeY = ShapeInference({fShapeX, fShapeW})[0];
333 if (!model.CheckIfTensorAlreadyExist(fNY)) {
334 model.AddIntermediateTensor(fNY, model.GetTensorType(fNX), fShapeY);
335 }
336 }
337 if (!fNY_h.empty()) {
338 fShapeY_h = ShapeInference({fShapeX, fShapeW})[1];
339 if (!model.CheckIfTensorAlreadyExist(fNY_h)) {
340 model.AddIntermediateTensor(fNY_h, model.GetTensorType(fNX), fShapeY_h);
341 }
342 }
343 if (!fNY_c.empty()) {
344 fShapeY_c = ShapeInference({fShapeX, fShapeW})[2];
345 if (!model.CheckIfTensorAlreadyExist(fNY_c)) {
346 model.AddIntermediateTensor(fNY_c, model.GetTensorType(fNX), fShapeY_c);
347 }
348 }
349 // Check the attributes
350 for (auto &activation : fAttrActivations) {
351 if (activation != "Relu" && activation != "Tanh" && activation != "Sigmoid" && activation != "Affine" &&
352 activation != "LeakyRelu" && activation != "ThresholdRelu" && activation != "ScaledTanh" &&
353 activation != "HardSigmoid" && activation != "Elu" && activation != "Softsign" && activation != "Softplus") {
354 throw std::runtime_error("TMVA SOFIE - Activation function " + activation + " not implemented");
355 }
356 }
357 if (fAttrDirection != "forward" && fAttrDirection != "backward" && fAttrDirection != "bidirectional") {
358 throw std::runtime_error("TMVA SOFIE - Invalid LSTM direction fAttrDirection = " + fAttrDirection);
359 }
360 if (4 * fAttrHiddenSize != fShapeW[1]) {
361 throw std::runtime_error("TMVA SOFIE - fAttrHiddenSize must be equal to " + std::to_string(fShapeW[1] / 4));
362 }
363 if (fAttrInputForget > 1) {
364 throw std::runtime_error("TMVA SOFIE - fAttrInputForget = " + std::to_string(fAttrInputForget) +
365 " must be 0 or 1.");
366 }
367 if (fAttrLayout > 1) {
368 throw std::runtime_error("TMVA SOFIE - Layout fAttrLayout = " + std::to_string(fAttrLayout) +
369 " must be 0 (timewise) or 1 (batchwise)");
370 }
371 if (fAttrActivations.empty()) {
372 if (fAttrDirection == "bidirectional") {
373 fAttrActivations = {"Sigmoid", "Tanh", "Tanh", "Sigmoid", "Tanh", "Tanh"};
374 } else {
375 fAttrActivations = {"Sigmoid", "Tanh", "Tanh"};
376 }
377 }
378}
379
380// generate code for Session data members (e.g. internal vectors)
381template <typename T>
383{
384 opName = "op_" + opName;
385 std::stringstream out;
386
387 size_t num_directions = fShapeW[0];
388 size_t seq_length = (fAttrLayout == 0) ? fShapeX[0] : fShapeX[1];
389 size_t batch_size = (fAttrLayout == 0) ? fShapeX[1] : fShapeX[0];
390 size_t input_size = fShapeX[2];
391
392 struct Block {
393 std::string name;
394 size_t size;
395 };
396
397 std::vector<Block> blocks;
398
399 size_t ff_size = seq_length * batch_size * fAttrHiddenSize;
400 size_t hs_size = seq_length * num_directions * batch_size * fAttrHiddenSize;
401
402 // Layout-dependent buffers
403 if (fAttrLayout != 0) {
404 blocks.push_back({"input", seq_length * batch_size * input_size});
405 blocks.push_back({"initial_hidden_state", num_directions * batch_size * fAttrHiddenSize});
406 blocks.push_back({"initial_cell_state", num_directions * batch_size * fAttrHiddenSize});
407 }
408
409 // Feedforward gates
410 blocks.push_back({"ff_input_gate", ff_size});
411 blocks.push_back({"ff_output_gate", ff_size});
412 blocks.push_back({"ff_cell_gate", ff_size});
413 if (fAttrInputForget == 0)
414 blocks.push_back({"ff_forget_gate", ff_size});
415
416 // Gate outputs
417 blocks.push_back({"input_gate", hs_size});
418 blocks.push_back({"output_gate", hs_size});
419 blocks.push_back({"cell_gate", hs_size});
420 if (fAttrInputForget == 0)
421 blocks.push_back({"forget_gate", hs_size});
422
423 // Cell state
424 blocks.push_back({"cell_state", hs_size});
425 blocks.push_back({"new_cell_state", hs_size});
426
427 // Hidden state (conditional)
428 if (fAttrLayout != 0 || fNY.empty()) {
429 blocks.push_back({"hidden_state", hs_size});
430 }
431
432 // Compute total size
433 size_t total_size = 0;
434 for (const auto &b : blocks) {
435 total_size += b.size;
436 }
437
438 // Backing storage
439 out << "std::vector<" << fType << "> fVec_" << opName << "_buffer = std::vector<" << fType << ">(" << total_size
440 << ");\n";
441
442 // Emit pointers
443 std::size_t offset = 0;
444 for (const auto &b : blocks) {
445 out << fType << "* fVec_" << opName << "_" << b.name << " = fVec_" << opName << "_buffer.data() + " << offset
446 << ";\n";
447 offset += b.size;
448 }
449
450 out << "\n";
451
452 return out.str();
453}
454
455template <typename T>
456auto ROperator_LSTM<T>::Generate(std::string OpName) -> std::string
457{
458 OpName = "op_" + OpName;
459 std::stringstream out;
460
461 size_t seq_length = (fAttrLayout == 0) ? fShapeX[0] : fShapeX[1];
462 size_t batch_size = (fAttrLayout == 0) ? fShapeX[1] : fShapeX[0];
463 size_t input_size = fShapeX[2];
464 size_t num_directions = fShapeW[0];
465
466 // set the input
467 if (fAttrLayout == 0) {
468 out << SP << fType << " const *" << OpName << "_input = tensor_" << fNX << ";\n";
469 } else {
470 out << SP << fType << " * " << OpName << "_input = this->fVec_" << OpName << "_input;\n";
471
472 out << SP << "for(size_t seq = 0; seq < " << seq_length << "; seq++) {\n";
473 out << SP << SP << "for(size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
474 out << SP << SP << SP << "for(size_t i = 0; i < " << input_size << "; i++) {\n";
475 out << SP << SP << SP << SP << OpName << "_input[seq * " << batch_size * input_size << " + batch * " << input_size
476 << " + i] = " << "tensor_" << fNX << "[batch * " << seq_length * input_size << " + seq * " << input_size
477 << " + i];\n";
478 out << SP << SP << SP << "}\n";
479 out << SP << SP << "}\n";
480 out << SP << "}\n";
481 }
482
483 // Set the initial hidden state
484 if (!fNInitial_h.empty()) {
485 if (fAttrLayout == 0) {
486 out << SP << fType << " const*" << OpName << "_initial_hidden_state = " << " tensor_" << fNInitial_h << ";\n";
487 } else {
488 out << SP << fType << " const* " << OpName << "_initial_hidden_state = this->fVec_" << OpName
489 << "_initial_hidden_state;\n";
490
491 for (size_t direction = 0; direction < num_directions; direction++) {
492 out << SP << "for(size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
493 out << SP << SP << "for(size_t h = 0; h < " << fAttrHiddenSize << "; h++) {\n";
494 out << SP << SP << SP << OpName << "_initial_hidden_state[" << direction * batch_size * fAttrHiddenSize
495 << " + batch * " << fAttrHiddenSize << " + h] = tensor_" << fNInitial_h << "[batch * "
496 << num_directions * fAttrHiddenSize << " + " << direction * fAttrHiddenSize << " + h];\n";
497 out << SP << SP << "}\n";
498 out << SP << "}\n";
499 }
500 }
501 }
502
503 // Set the initial cell state
504 if (!fNInitial_c.empty()) {
505 if (fAttrLayout == 0) {
506 out << SP << fType << " const*" << OpName << "_initial_cell_state = " << " tensor_" << fNInitial_c << ";\n";
507 } else {
508 out << SP << fType << " const* " << OpName << "_initial_cell_state = this->fVec_" << OpName
509 << "_initial_cell_state;\n";
510
511 for (size_t direction = 0; direction < num_directions; direction++) {
512 out << SP << "for(size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
513 out << SP << SP << "for(size_t h = 0; h < " << fAttrHiddenSize << "; h++) {\n";
514 out << SP << SP << SP << OpName << "_initial_cell_state[" << direction * batch_size * fAttrHiddenSize
515 << " + batch * " << fAttrHiddenSize << " + h] = tensor_" << fNInitial_c << "[batch * "
516 << num_directions * fAttrHiddenSize << " + " << direction * fAttrHiddenSize << " + h];\n";
517 out << SP << SP << "}\n";
518 out << SP << "}\n";
519 }
520 }
521 }
522
523 // Set the feedforward
524 out << SP << fType << " * " << OpName << "_ff_input_gate = this->fVec_" << OpName << "_ff_input_gate;\n";
525 out << SP << fType << " * " << OpName << "_ff_output_gate = this->fVec_" << OpName << "_ff_output_gate;\n";
526 out << SP << fType << " * " << OpName << "_ff_cell_gate = this->fVec_" << OpName << "_ff_cell_gate;\n";
527 if (fAttrInputForget == 0) {
528 out << SP << fType << " * " << OpName << "_ff_forget_gate = this->fVec_" << OpName << "_ff_forget_gate;\n";
529 }
530 // Set the gates
531 out << SP << fType << " * " << OpName << "_input_gate = this->fVec_" << OpName << "_input_gate;\n";
532 out << SP << fType << " * " << OpName << "_output_gate = this->fVec_" << OpName << "_output_gate;\n";
533 out << SP << fType << " * " << OpName << "_cell_gate = this->fVec_" << OpName << "_cell_gate;\n";
534 if (fAttrInputForget == 0) {
535 out << SP << fType << " * " << OpName << "_forget_gate = this->fVec_" << OpName << "_forget_gate;\n";
536 }
537 // Set the cell state and the new cell state = h(cell state)
538 out << SP << fType << " * " << OpName << "_cell_state = this->fVec_" << OpName << "_cell_state;\n";
539 out << SP << fType << " * " << OpName << "_new_cell_state = this->fVec_" << OpName << "_new_cell_state;\n";
540
541 // Set the hidden state
542 if (fAttrLayout == 0 && !fNY.empty()) {
543 out << SP << fType << " *" << OpName << "_hidden_state = tensor_" << fNY << ";\n";
544 } else {
545 out << SP << fType << " * " << OpName << "_hidden_state = this->fVec_" << OpName << "_hidden_state;\n";
546 }
547
548 out << SP << "char " << OpName << "_transA = 'N';\n";
549 out << SP << "char " << OpName << "_transB = 'T';\n";
550 out << SP << "int " << OpName << "_m = " << seq_length * batch_size << ";\n";
551 out << SP << "int " << OpName << "_n = " << fAttrHiddenSize << ";\n";
552 out << SP << "int " << OpName << "_k = " << input_size << ";\n";
553 if (fType == "float") {
554 out << SP << fType << " " << OpName << "_alpha = 1.;\n";
555 out << SP << fType << " " << OpName << "_beta = 0.;\n";
556 }
557 if (!fNB.empty()) {
558 out << SP << "int " << OpName << "_bias_size = " << seq_length * batch_size * fAttrHiddenSize << ";\n";
559 out << SP << "int " << OpName << "_incx = 1;\n";
560 out << SP << "int " << OpName << "_incy = 1;\n";
561 }
562
563 auto emit_sgemm = [&](const std::string &out_name, size_t offset) -> std::string {
564 std::stringstream ss;
565 ss << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &" << OpName << "_n, &" << OpName
566 << "_m, &" << OpName << "_k, &" << OpName << "_alpha, tensor_" << fNW;
567
568 if (offset != 0)
569 ss << " + " << offset;
570
571 ss << ", &" << OpName << "_k, " << OpName << "_input, &" << OpName << "_k, &" << OpName << "_beta, " << OpName
572 << "_" << out_name << ", &" << OpName << "_n);\n";
573 return ss.str();
574 };
575
576 for (size_t direction = 0; direction < num_directions; direction++) {
577 if (direction == 0) {
578 if (fType == "float") {
579 // input_gate = input * weight_i^T
580 out << SP << emit_sgemm("ff_input_gate", 0);
581 // output_gate = input * weight_o^T
582 size_t wo_offset = fAttrHiddenSize * input_size;
583 out << SP << emit_sgemm("ff_output_gate", wo_offset);
584 // cell_gate = input * weight_c^T
585 size_t wc_offset = 3 * fAttrHiddenSize * input_size;
586 out << SP << emit_sgemm("ff_cell_gate", wc_offset);
587 }
588 } else {
589 if (fType == "float") {
590 // input_gate = input * weight_i^T
591 out << SP << emit_sgemm("ff_input_gate", 4 * fAttrHiddenSize * input_size);
592 // output_gate = input * weight_o^T
593 size_t wo_offset = 4 * fAttrHiddenSize * input_size + 1 * fAttrHiddenSize * input_size;
594 out << SP << emit_sgemm("ff_output_gate", wo_offset);
595 // cell_gate = input * weight_c^T
596 size_t wc_offset = 4 * fAttrHiddenSize * input_size + 3 * fAttrHiddenSize * input_size;
597 out << SP << emit_sgemm("ff_cell_gate", wc_offset);
598 }
599 }
600 if (fAttrInputForget == 0) {
601 // forget_gate = input * weight_f^T
602 if (direction == 0) {
603 if (fType == "float") {
604 size_t wf_offset = 2 * fAttrHiddenSize * input_size;
605 out << SP << emit_sgemm("ff_forget_gate", wf_offset);
606 }
607 } else {
608 if (fType == "float") {
609 size_t wf_offset = 4 * fAttrHiddenSize * input_size + 2 * fAttrHiddenSize * input_size;
610 out << SP << emit_sgemm("ff_forget_gate", wf_offset);
611 }
612 }
613 }
614
615 // Add the bias
616 if (!fNB.empty()) {
617 if (direction == 0) {
618 if (fType == "float") {
619 // ff_input_gate += bias_i
620 out << SP << "BLAS::saxpy_(&" << OpName << "_bias_size, &" << OpName << "_alpha, tensor_" << fNB << ", &"
621 << OpName << "_incx, " << OpName << "_ff_input_gate, &" << OpName << "_incy);\n";
622 // ff_output_gate += bias_o
623 size_t bo_offset = seq_length * batch_size * fAttrHiddenSize;
624 out << SP << "BLAS::saxpy_(&" << OpName << "_bias_size, &" << OpName << "_alpha, tensor_" << fNB << " + "
625 << bo_offset << ", &" << OpName << "_incx, " << OpName << "_ff_output_gate, &" << OpName
626 << "_incy);\n";
627 // ff_cell_gate += bias_c
628 size_t bc_offset = 3 * seq_length * batch_size * fAttrHiddenSize;
629 out << SP << "BLAS::saxpy_(&" << OpName << "_bias_size, &" << OpName << "_alpha, tensor_" << fNB << " + "
630 << bc_offset << ", &" << OpName << "_incx, " << OpName << "_ff_cell_gate, &" << OpName
631 << "_incy);\n";
632 }
633 } else {
634 if (fType == "float") {
635 // ff_input_gate += bias_i
636 size_t bi_offset = 4 * seq_length * batch_size * fAttrHiddenSize;
637 out << SP << "BLAS::saxpy_(&" << OpName << "_bias_size, &" << OpName << "_alpha, tensor_" << fNB << " + "
638 << bi_offset << ", &" << OpName << "_incx, " << OpName << "_ff_input_gate, &" << OpName
639 << "_incy);\n";
640 // ff_output_gate += bias_o
641 size_t bo_offset =
642 4 * seq_length * batch_size * fAttrHiddenSize + seq_length * batch_size * fAttrHiddenSize;
643 out << SP << "BLAS::saxpy_(&" << OpName << "_bias_size, &" << OpName << "_alpha, tensor_" << fNB << " + "
644 << bo_offset << ", &" << OpName << "_incx, " << OpName << "_ff_output_gate, &" << OpName
645 << "_incy);\n";
646 // ff_cell_gate += bias_c
647 size_t bc_offset = 4 * num_directions * seq_length * batch_size * fAttrHiddenSize +
648 3 * seq_length * batch_size * fAttrHiddenSize;
649 out << SP << "BLAS::saxpy_(&" << OpName << "_bias_size, &" << OpName << "_alpha, tensor_" << fNB << " + "
650 << bc_offset << ", &" << OpName << "_incx, " << OpName << "_ff_cell_gate, &" << OpName
651 << "_incy);\n";
652 }
653 }
654 if (fAttrInputForget == 0) {
655 // ff_forget_gate += bias_f
656 if (direction == 0) {
657 if (fType == "float") {
658 size_t bo_offset = 2 * seq_length * batch_size * fAttrHiddenSize;
659 out << SP << "BLAS::saxpy_(&" << OpName << "_bias_size, &" << OpName << "_alpha, tensor_" << fNB
660 << " + " << bo_offset << ", &" << OpName << "_incx, " << OpName << "_ff_forget_gate, &" << OpName
661 << "_incy);\n";
662 }
663 } else {
664 if (fType == "float") {
665 size_t bo_offset =
666 4 * seq_length * batch_size * fAttrHiddenSize + 2 * seq_length * batch_size * fAttrHiddenSize;
667 out << SP << "BLAS::saxpy_(&" << OpName << "_bias_size, &" << OpName << "_alpha, tensor_" << fNB
668 << " + " << bo_offset << ", &" << OpName << "_incx, " << OpName << "_ff_forget_gate, &" << OpName
669 << "_incy);\n";
670 }
671 }
672 }
673 }
674
675 // Copy ff_input_gate, ff_output_gate, ff_cell_gate and ff_forget_gate into input_gate, output_gate,
676 // cell_gate and forget_gate
677 out << SP << "for (size_t seq = 0; seq < " << seq_length << "; seq++) {\n";
678 out << SP << SP << "size_t ff_offset = seq * " << batch_size * fAttrHiddenSize << ";\n";
679 if (direction == 0) {
680 out << SP << SP << "size_t gate_offset = seq * " << num_directions * batch_size * fAttrHiddenSize << ";\n";
681 } else {
682 out << SP << SP << "size_t gate_offset = seq * " << num_directions * batch_size * fAttrHiddenSize << " + "
683 << batch_size * fAttrHiddenSize << ";\n";
684 }
685 size_t ff_seq_size = batch_size * fAttrHiddenSize;
686 out << SP << SP << "std::copy(" << OpName << "_ff_input_gate + ff_offset, " << OpName
687 << "_ff_input_gate + ff_offset + " << ff_seq_size << ", " << OpName << "_input_gate + gate_offset);\n";
688 out << SP << SP << "std::copy(" << OpName << "_ff_output_gate + ff_offset, " << OpName
689 << "_ff_output_gate + ff_offset + " << ff_seq_size << ", " << OpName << "_output_gate + gate_offset);\n";
690 out << SP << SP << "std::copy(" << OpName << "_ff_cell_gate + ff_offset, " << OpName
691 << "_ff_cell_gate + ff_offset + " << ff_seq_size << ", " << OpName << "_cell_gate + gate_offset);\n";
692 if (fAttrInputForget == 0) {
693 out << SP << SP << "std::copy(" << OpName << "_ff_forget_gate + ff_offset, " << OpName
694 << "_ff_forget_gate + ff_offset + " << ff_seq_size << ", " << OpName << "_forget_gate + gate_offset);\n";
695 }
696 out << SP << "}\n";
697
698 out << SP << "for (size_t seq = 0; seq < " << seq_length << "; seq++) {\n";
699 if (fAttrDirection == "backward" || direction == 1) {
700 out << SP << SP << "size_t index = " << seq_length - 1 << " - seq;\n";
701 } else {
702 out << SP << SP << "size_t index = seq;\n";
703 }
704 out << SP << SP << "int m2 = " << batch_size << ";\n";
705 if (direction == 0) {
706 out << SP << SP << "size_t offset = index * " << num_directions * batch_size * fAttrHiddenSize << ";\n";
707 } else {
708 out << SP << SP << "size_t offset = index * " << num_directions * batch_size * fAttrHiddenSize << " + "
709 << batch_size * fAttrHiddenSize << ";\n";
710 }
711 size_t size = batch_size * fAttrHiddenSize;
712 // gate = gate + initial_hidden_state * Recurrence^T
713 out << SP << SP << "if (seq == 0) {\n";
714 if (!fNInitial_h.empty()) {
715 if (direction == 0) {
716 if (fType == "float") {
717 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &" << OpName
718 << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << ", &" << OpName
719 << "_n, " << OpName << "_initial_hidden_state, &" << OpName << "_n, &" << OpName << "_alpha, "
720 << OpName << "_input_gate + offset, &" << OpName << "_n);\n";
721 size_t ro_offset = fAttrHiddenSize * fAttrHiddenSize;
722 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &" << OpName
723 << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << " + " << ro_offset
724 << ", &" << OpName << "_n, " << OpName << "_initial_hidden_state, &" << OpName << "_n, &" << OpName
725 << "_alpha, " << OpName << "_output_gate + offset, &" << OpName << "_n);\n";
726 size_t rc_offset = 3 * fAttrHiddenSize * fAttrHiddenSize;
727 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &" << OpName
728 << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << " + " << rc_offset
729 << ", &" << OpName << "_n, " << OpName << "_initial_hidden_state, &" << OpName << "_n, &" << OpName
730 << "_alpha, " << OpName << "_cell_gate + offset, &" << OpName << "_n);\n";
731 if (fAttrInputForget == 0) {
732 size_t rf_offset = 2 * fAttrHiddenSize * fAttrHiddenSize;
733 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &"
734 << OpName << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << " + "
735 << rf_offset << ", &" << OpName << "_n, " << OpName << "_initial_hidden_state, &" << OpName
736 << "_n, &" << OpName << "_alpha, " << OpName << "_forget_gate + offset, &" << OpName << "_n);\n";
737 }
738 }
739 } else { // direction=1
740 if (fType == "float") {
741 size_t ri_offset = 4 * fAttrHiddenSize * fAttrHiddenSize;
742 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &" << OpName
743 << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << " + " << ri_offset
744 << ", &" << OpName << "_n, " << OpName << "_initial_hidden_state, &" << OpName << "_n, &" << OpName
745 << "_alpha, " << OpName << "_input_gate + offset, &" << OpName << "_n);\n";
746 size_t ro_offset = 4 * fAttrHiddenSize * fAttrHiddenSize + 1 * fAttrHiddenSize * fAttrHiddenSize;
747 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &" << OpName
748 << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << " + " << ro_offset
749 << ", &" << OpName << "_n, " << OpName << "_initial_hidden_state, &" << OpName << "_n, &" << OpName
750 << "_alpha, " << OpName << "_output_gate + offset, &" << OpName << "_n);\n";
751 size_t rc_offset = 4 * fAttrHiddenSize * fAttrHiddenSize + 3 * fAttrHiddenSize * fAttrHiddenSize;
752 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &" << OpName
753 << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << " + " << rc_offset
754 << ", &" << OpName << "_n, " << OpName << "_initial_hidden_state, &" << OpName << "_n, &" << OpName
755 << "_alpha, " << OpName << "_cell_gate + offset, &" << OpName << "_n);\n";
756 if (fAttrInputForget == 0) {
757 size_t rf_offset = 4 * fAttrHiddenSize * fAttrHiddenSize + 2 * fAttrHiddenSize * fAttrHiddenSize;
758 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &"
759 << OpName << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << " + "
760 << rf_offset << ", &" << OpName << "_n, " << OpName << "_initial_hidden_state, &" << OpName
761 << "_n, &" << OpName << "_alpha, " << OpName << "_forget_gate + offset, &" << OpName << "_n);\n";
762 }
763 }
764 }
765 }
766 out << SP << SP << "} else {\n";
767 // gate = gate + previous_hidden_state * Recurrence^T
768 if (direction == 0) {
769 if (fAttrDirection == "backward") {
770 out << SP << SP << SP << "size_t previous_offset = (index + 1) * "
771 << num_directions * batch_size * fAttrHiddenSize << ";\n";
772 } else {
773 out << SP << SP << SP << "size_t previous_offset = (seq - 1) * "
774 << num_directions * batch_size * fAttrHiddenSize << ";\n";
775 }
776 if (fType == "float") {
777 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &" << OpName
778 << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << ", &" << OpName << "_n, "
779 << OpName << "_hidden_state + previous_offset, &" << OpName << "_n, &" << OpName << "_alpha, " << OpName
780 << "_input_gate + offset, &" << OpName << "_n);\n";
781 size_t ro_offset = 1 * fAttrHiddenSize * fAttrHiddenSize;
782 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &" << OpName
783 << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << " + " << ro_offset
784 << ", &" << OpName << "_n, " << OpName << "_hidden_state + previous_offset, &" << OpName << "_n, &"
785 << OpName << "_alpha, " << OpName << "_output_gate + offset, &" << OpName << "_n);\n";
786 size_t rc_offset = 3 * fAttrHiddenSize * fAttrHiddenSize;
787 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &" << OpName
788 << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << " + " << rc_offset
789 << ", &" << OpName << "_n, " << OpName << "_hidden_state + previous_offset, &" << OpName << "_n, &"
790 << OpName << "_alpha, " << OpName << "_cell_gate + offset, &" << OpName << "_n);\n";
791 if (fAttrInputForget == 0) {
792 size_t rf_offset = 2 * fAttrHiddenSize * fAttrHiddenSize;
793 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &" << OpName
794 << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << " + " << rf_offset
795 << ", &" << OpName << "_n, " << OpName << "_hidden_state + previous_offset, &" << OpName << "_n, &"
796 << OpName << "_alpha, " << OpName << "_forget_gate + offset, &" << OpName << "_n);\n";
797 }
798 }
799 } else {
800 out << SP << SP << SP << "size_t previous_offset = (index + 1) * "
801 << num_directions * batch_size * fAttrHiddenSize << " + " << batch_size * fAttrHiddenSize << ";\n";
802 if (fType == "float") {
803 size_t ri_offset = 4 * fAttrHiddenSize * fAttrHiddenSize;
804 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &" << OpName
805 << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << " + " << ri_offset
806 << ", &" << OpName << "_n, " << OpName << "_hidden_state + previous_offset, &" << OpName << "_n, &"
807 << OpName << "_alpha, " << OpName << "_input_gate + offset, &" << OpName << "_n);\n";
808 size_t ro_offset = 4 * fAttrHiddenSize * fAttrHiddenSize + fAttrHiddenSize * fAttrHiddenSize;
809 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &" << OpName
810 << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << " + " << ro_offset
811 << ", &" << OpName << "_n, " << OpName << "_hidden_state + previous_offset, &" << OpName << "_n, &"
812 << OpName << "_alpha, " << OpName << "_output_gate + offset, &" << OpName << "_n);\n";
813 size_t rc_offset = 4 * fAttrHiddenSize * fAttrHiddenSize + 3 * fAttrHiddenSize * fAttrHiddenSize;
814 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &" << OpName
815 << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << " + " << rc_offset
816 << ", &" << OpName << "_n, " << OpName << "_hidden_state + previous_offset, &" << OpName << "_n, &"
817 << OpName << "_alpha, " << OpName << "_cell_gate + offset, &" << OpName << "_n);\n";
818 if (fAttrInputForget == 0) {
819 size_t rf_offset = 4 * fAttrHiddenSize * fAttrHiddenSize + 2 * fAttrHiddenSize * fAttrHiddenSize;
820 out << SP << SP << SP << "BLAS::sgemm_(&" << OpName << "_transB, &" << OpName << "_transA, &" << OpName
821 << "_n, &m2, &" << OpName << "_n, &" << OpName << "_alpha, tensor_" << fNR << " + " << rf_offset
822 << ", &" << OpName << "_n, " << OpName << "_hidden_state + previous_offset, &" << OpName << "_n, &"
823 << OpName << "_alpha, " << OpName << "_forget_gate + offset, &" << OpName << "_n);\n";
824 }
825 }
826 }
827 out << SP << SP << "}\n";
828
829 // Clip the elements of the cell gate into the range [-fAttrClip, fAttrClip]
830 if (fAttrClip > .0) {
831 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
832 if (fType == "float") {
833 out << SP << SP << SP << "float x = (" << OpName << "_cell_gate[i] > " << -fAttrClip << ") ? " << OpName
834 << "_cell_gate[i] : " << -fAttrClip << ";\n";
835 }
836 out << SP << SP << SP << OpName << "_cell_gate[i] = (x < " << fAttrClip << ") ? x : " << fAttrClip << ";\n";
837 out << SP << SP << "}\n";
838 }
839 // Apply the activation function to the cell gate, cell_gate = g(cell_gate)
840 if (fAttrActivations[direction * 3 + 1] == "Relu") {
841 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
842 out << SP << SP << SP << "if (" << OpName << "_cell_gate[i] < 0.)\n";
843 out << SP << SP << SP << SP << OpName << "_cell_gate[i] = 0.;\n";
844 out << SP << SP << "}\n";
845 } else if (fAttrActivations[direction * 3 + 1] == "Tanh") {
846 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
847 if (fType == "float") {
848 out << SP << SP << SP << "float ex = exp(-2 * " << OpName << "_cell_gate[i]);\n";
849 }
850 out << SP << SP << SP << SP << OpName << "_cell_gate[i] = (1. - ex) / (1. + ex);\n";
851 out << SP << SP << "}\n";
852 } else if (fAttrActivations[direction * 3 + 1] == "Sigmoid") {
853 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
854 out << SP << SP << SP << SP << OpName << "_cell_gate[i] = 1. / (1. + exp(-" << OpName << "_cell_gate[i]));\n";
855 out << SP << SP << "}\n";
856 } else if (fAttrActivations[direction * 3 + 1] == "Affine") {
857 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
858 out << SP << SP << SP << SP << OpName << "_cell_gate[i] = " << fAttrActivationAlpha[direction * 3 + 1] << " * "
859 << OpName << "_cell_gate[i] + " << fAttrActivationBeta[direction * 3 + 1] << ";\n";
860 out << SP << SP << "}\n";
861 } else if (fAttrActivations[direction * 3 + 1] == "ScaledTanh") {
862 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
863 if (fType == "float") {
864 out << SP << SP << SP << "float ex = exp(-2 * " << fAttrActivationBeta[direction * 3 + 1] << " * " << OpName
865 << "_cell_gate[i]);\n";
866 }
867 out << SP << SP << SP << SP << OpName << "_cell_gate[i] = " << fAttrActivationAlpha[direction * 3 + 1]
868 << " * (1. - ex) / (1. + ex);\n";
869 out << SP << SP << "}\n";
870 } else if (fAttrActivations[direction * 3 + 1] == "HardSigmoid") {
871 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
872 if (fType == "float") {
873 out << SP << SP << SP << "float a = " << fAttrActivationAlpha[direction * 3 + 1] << " * " << OpName
874 << "_cell_gate[i] + " << fAttrActivationBeta[direction * 3 + 1] << ";\n";
875 out << SP << SP << SP << "float b = (a > 0.) ? a : 0.;\n";
876 }
877 out << SP << SP << SP << SP << OpName << "_cell_gate[i] = (b < 1.) ? b : 1.;\n";
878 out << SP << SP << "}\n";
879 } else if (fAttrActivations[direction * 3 + 1] == "LeakyRelu") {
880 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
881 out << SP << SP << SP << "if (" << OpName << "_cell_gate[i] < 0.)\n";
882 out << SP << SP << SP << SP << OpName << "_cell_gate[i] = " << fAttrActivationAlpha[direction * 3 + 1] << " * "
883 << OpName << "_cell_gate[i];\n";
884 out << SP << SP << "}\n";
885 } else if (fAttrActivations[direction * 3 + 1] == "ThresholdRelu") {
886 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
887 out << SP << SP << SP << "if (" << OpName << "_cell_gate[i] < " << fAttrActivationAlpha[direction * 3 + 1]
888 << ")\n";
889 out << SP << SP << SP << SP << OpName << "_cell_gate[i] = 0.;\n";
890 out << SP << SP << "}";
891 } else if (fAttrActivations[direction * 3 + 1] == "Elu") {
892 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
893 out << SP << SP << SP << "if (" << OpName << "_cell_gate[i] < 0.)\n";
894 out << SP << SP << SP << SP << OpName << "_cell_gate[i] = " << fAttrActivationAlpha[direction * 3 + 1]
895 << " * exp(" << OpName << "_cell_gate[i] - 1.);\n";
896 out << SP << SP << "}\n";
897 } else if (fAttrActivations[direction * 3 + 1] == "Softsign") {
898 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
899 out << SP << SP << SP << SP << OpName << "_cell_gate[i] = " << OpName << "_cell_gate[i] / (1. + abs(" << OpName
900 << "_cell_gate[i]));\n";
901 out << SP << SP << "}\n";
902 } else { // fAttrActivations[direction * 3 + 1] = Softplus
903 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
904 out << SP << SP << SP << SP << OpName << "_cell_gate[i] = log(1. + exp(" << OpName << "_cell_gate[i]));\n";
905 out << SP << SP << "}\n";
906 }
907
908 // Peephole connections for the input gate and the forget gate
909 if (!fNP.empty()) {
910 // gate = 1.0 * gate + previous_cell_state * P^T
911 out << SP << SP << "if (seq == 0) {\n";
912 if (!fNInitial_c.empty()) {
913 if (direction == 0) {
914 out << SP << SP << SP << "for (size_t i = 0; i < " << size << "; i++) {\n";
915 out << SP << SP << SP << SP << OpName << "_input_gate[i + offset] += tensor_" << fNP << "[i] * "
916 << OpName << "_initial_cell_state[i];\n";
917 out << SP << SP << SP << "}\n";
918 if (fAttrInputForget == 0) {
919 size_t pf_offset = batch_size * fAttrHiddenSize;
920 out << SP << SP << SP << "for (size_t i = 0; i < " << size << "; i++) {\n";
921 out << SP << SP << SP << SP << OpName << "_forget_gate[i + offset] += tensor_" << fNP << "[i + "
922 << pf_offset << "] * " << OpName << "_initial_cell_state[i];\n";
923 out << SP << SP << SP << "}\n";
924 }
925 } else {
926 size_t pi_offset = 3 * batch_size * fAttrHiddenSize;
927 size_t initial_c_offset = batch_size * fAttrHiddenSize;
928 out << SP << SP << SP << "for (size_t i = 0; i < " << size << "; i++) {\n";
929 out << SP << SP << SP << SP << OpName << "_input_gate[i + offset] += tensor_" << fNP << "[i + "
930 << pi_offset << "] * " << OpName << "_initial_cell_state[i + " << initial_c_offset << "];\n";
931 out << SP << SP << SP << "}\n";
932 if (fAttrInputForget == 0) {
933 size_t pf_offset = 3 * batch_size * fAttrHiddenSize + batch_size * fAttrHiddenSize;
934 out << SP << SP << SP << "for (size_t i = 0; i < " << size << "; i++) {\n";
935 out << SP << SP << SP << SP << OpName << "_forget_gate[i + offset] += tensor_" << fNP << "[i + "
936 << pf_offset << "] * " << OpName << "_initial_cell_state[i + " << initial_c_offset << "];\n";
937 out << SP << SP << SP << "}\n";
938 }
939 }
940 }
941 out << SP << SP << "} else {\n";
942 if (direction == 0) {
943 if (fAttrDirection == "backward") {
944 out << SP << SP << SP << "size_t c_offset = (index + 1) * "
945 << num_directions * batch_size * fAttrHiddenSize << ";\n";
946 } else {
947 out << SP << SP << SP << "size_t c_offset = (seq - 1) * "
948 << num_directions * batch_size * fAttrHiddenSize << ";\n";
949 }
950 out << SP << SP << SP << "for (size_t i = 0; i < " << size << "; i++) {\n";
951 out << SP << SP << SP << SP << OpName << "_input_gate[i + offset] += tensor_" << fNP << "[i] * " << OpName
952 << "_cell_state[i + c_offset];\n";
953 out << SP << SP << SP << "}\n";
954 if (fAttrInputForget == 0) {
955 size_t pf_offset = batch_size * fAttrHiddenSize;
956 out << SP << SP << SP << "for (size_t i = 0; i < " << size << "; i++) {\n";
957 out << SP << SP << SP << SP << OpName << "_forget_gate[i + offset] += tensor_" << fNP << "[i + "
958 << pf_offset << "] * " << OpName << "_cell_state[i + c_offset];\n";
959 out << SP << SP << SP << "}\n";
960 }
961 } else { // direction=1
962 size_t pi_offset = 3 * batch_size * fAttrHiddenSize;
963 out << SP << SP << SP << "size_t c_offset = (index + 1) * " << num_directions * batch_size * fAttrHiddenSize
964 << " + " << batch_size * fAttrHiddenSize << ";\n";
965 out << SP << SP << SP << "for (size_t i = 0; i < " << size << "; i++) {\n";
966 out << SP << SP << SP << SP << OpName << "_input_gate[i + offset] += tensor_" << fNP << "[i + " << pi_offset
967 << "] * " << OpName << "_cell_state[i + c_offset];\n";
968 out << SP << SP << SP << "}\n";
969 if (fAttrInputForget == 0) {
970 size_t pf_offset = 3 * batch_size * fAttrHiddenSize + batch_size * fAttrHiddenSize;
971 out << SP << SP << SP << "for (size_t i = 0; i < " << size << "; i++) {\n";
972 out << SP << SP << SP << SP << OpName << "_forget_gate[i + offset] += tensor_" << fNP << "[i + "
973 << pf_offset << "] * " << OpName << "_cell_state[i + c_offset];\n";
974 out << SP << SP << SP << "}\n";
975 }
976 }
977 out << SP << SP << "}\n";
978 }
979
980 // Clip the elements of the input gate into the range [-fAttrClip, fAttrClip]
981 if (fAttrClip > .0) {
982 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
983 if (fType == "float") {
984 out << SP << SP << SP << "float x = (" << OpName << "_input_gate[i] > " << -fAttrClip << ") ? " << OpName
985 << "_input_gate[i] : " << -fAttrClip << ";\n";
986 }
987 out << SP << SP << SP << OpName << "_input_gate[i] = (x < " << fAttrClip << ") ? x : " << fAttrClip << ";\n";
988 out << SP << SP << "}\n";
989 }
990 // Apply the activation function to the input gate
991 if (fAttrActivations[direction * 3] == "Relu") {
992 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
993 out << SP << SP << SP << "if (" << OpName << "_input_gate[i] < 0.)\n";
994 out << SP << SP << SP << SP << OpName << "_input_gate[i] = 0.;\n";
995 out << SP << SP << "}\n";
996 } else if (fAttrActivations[direction * 3] == "Tanh") {
997 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
998 if (fType == "float") {
999 out << SP << SP << SP << "float ex = exp(-2 * " << OpName << "_input_gate[i]);\n";
1000 }
1001 out << SP << SP << SP << SP << OpName << "_input_gate[i] = (1. - ex) / (1. + ex);\n";
1002 out << SP << SP << "}\n";
1003 } else if (fAttrActivations[direction * 3] == "Sigmoid") {
1004 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1005 out << SP << SP << SP << SP << OpName << "_input_gate[i] = 1. / (1. + exp(-" << OpName
1006 << "_input_gate[i]));\n";
1007 out << SP << SP << "}\n";
1008 } else if (fAttrActivations[direction * 3] == "Affine") {
1009 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1010 out << SP << SP << SP << SP << OpName << "_input_gate[i] = " << fAttrActivationAlpha[direction * 3] << " * "
1011 << OpName << "_input_gate[i] + " << fAttrActivationBeta[direction * 3] << ";\n";
1012 out << SP << SP << "}\n";
1013 } else if (fAttrActivations[direction * 3] == "ScaledTanh") {
1014 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1015 if (fType == "float") {
1016 out << SP << SP << SP << "float ex = exp(-2 * " << fAttrActivationBeta[direction * 3] << " * " << OpName
1017 << "_input_gate[i]);\n";
1018 }
1019 out << SP << SP << SP << SP << OpName << "_input_gate[i] = " << fAttrActivationAlpha[direction * 3]
1020 << " * (1. - ex) / (1. + ex);\n";
1021 out << SP << SP << "}\n";
1022 } else if (fAttrActivations[direction * 3] == "HardSigmoid") {
1023 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1024 if (fType == "float") {
1025 out << SP << SP << SP << "float a = " << fAttrActivationAlpha[direction * 3] << " * " << OpName
1026 << "_input_gate[i] + " << fAttrActivationBeta[direction * 3] << ";\n";
1027 out << SP << SP << SP << "float b = (a > 0.) ? a : 0.;\n";
1028 }
1029 out << SP << SP << SP << SP << OpName << "_input_gate[i] = (b < 1.) ? b : 1.;\n";
1030 out << SP << SP << "}\n";
1031 } else if (fAttrActivations[direction * 3] == "LeakyRelu") {
1032 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1033 out << SP << SP << SP << "if (" << OpName << "_input_gate[i] < 0.)\n";
1034 out << SP << SP << SP << SP << OpName << "_input_gate[i] = " << fAttrActivationAlpha[direction * 3] << " * "
1035 << OpName << "_input_gate[i];\n";
1036 out << SP << SP << "}\n";
1037 } else if (fAttrActivations[direction * 3] == "ThresholdRelu") {
1038 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1039 out << SP << SP << SP << "if (" << OpName << "_input_gate[i] < " << fAttrActivationAlpha[direction * 3]
1040 << ")\n";
1041 out << SP << SP << SP << SP << OpName << "_input_gate[i] = 0.;\n";
1042 out << SP << SP << "}";
1043 } else if (fAttrActivations[direction * 3] == "Elu") {
1044 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1045 out << SP << SP << SP << "if (" << OpName << "_input_gate[i] < 0.)\n";
1046 out << SP << SP << SP << SP << OpName << "_input_gate[i] = " << fAttrActivationAlpha[direction * 3]
1047 << " * exp(" << OpName << "_input_gate[i] - 1.);\n";
1048 out << SP << SP << "}\n";
1049 } else if (fAttrActivations[direction * 3] == "Softsign") {
1050 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1051 out << SP << SP << SP << SP << OpName << "_input_gate[i] = " << OpName << "_input_gate[i] / (1. + abs("
1052 << OpName << "_input_gate[i]));\n";
1053 out << SP << SP << "}\n";
1054 } else { // fAttrActivations[direction * 3] = Softplus
1055 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1056 out << SP << SP << SP << SP << OpName << "_input_gate[i] = log(1. + exp(" << OpName << "_input_gate[i]));\n";
1057 out << SP << SP << "}\n";
1058 }
1059
1060 if (fAttrInputForget == 0) {
1061 // Clip the elements of the forget gate into the range [-fAttrClip, fAttrClip]
1062 if (fAttrClip > .0) {
1063 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1064 if (fType == "float") {
1065 out << SP << SP << SP << "float x = (" << OpName << "_forget_gate[i] > " << -fAttrClip << ") ? "
1066 << OpName << "_forget_gate[i] : " << -fAttrClip << ";\n";
1067 }
1068 out << SP << SP << SP << OpName << "_forget_gate[i] = (x < " << fAttrClip << ") ? x : " << fAttrClip
1069 << ";\n";
1070 out << SP << SP << "}\n";
1071 }
1072 // Apply the activation function to the forget gate, cell_gate = g(cell_gate)
1073 if (fAttrActivations[direction * 3] == "Relu") {
1074 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1075 out << SP << SP << SP << "if (" << OpName << "_forget_gate[i] < 0.)\n";
1076 out << SP << SP << SP << SP << OpName << "_forget_gate[i] = 0.;\n";
1077 out << SP << SP << "}\n";
1078 } else if (fAttrActivations[direction * 3] == "Tanh") {
1079 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1080 if (fType == "float") {
1081 out << SP << SP << SP << "float ex = exp(-2 * " << OpName << "_forget_gate[i]);\n";
1082 }
1083 out << SP << SP << SP << SP << OpName << "_forget_gate[i] = (1. - ex) / (1. + ex);\n";
1084 out << SP << SP << "}\n";
1085 } else if (fAttrActivations[direction * 3] == "Sigmoid") {
1086 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1087 out << SP << SP << SP << SP << OpName << "_forget_gate[i] = 1. / (1. + exp(-" << OpName
1088 << "_forget_gate[i]));\n";
1089 out << SP << SP << "}\n";
1090 } else if (fAttrActivations[direction * 3] == "Affine") {
1091 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1092 out << SP << SP << SP << SP << OpName << "_forget_gate[i] = " << fAttrActivationAlpha[direction * 3]
1093 << " * " << OpName << "_forget_gate[i] + " << fAttrActivationBeta[direction * 3] << ";\n";
1094 out << SP << SP << "}\n";
1095 } else if (fAttrActivations[direction * 3] == "ScaledTanh") {
1096 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1097 if (fType == "float") {
1098 out << SP << SP << SP << "float ex = exp(-2 * " << fAttrActivationBeta[direction * 3] << " * " << OpName
1099 << "_forget_gate[i]);\n";
1100 }
1101 out << SP << SP << SP << SP << OpName << "_forget_gate[i] = " << fAttrActivationAlpha[direction * 3]
1102 << " * (1. - ex) / (1. + ex);\n";
1103 out << SP << SP << "}\n";
1104 } else if (fAttrActivations[direction * 3] == "HardSigmoid") {
1105 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1106 if (fType == "float") {
1107 out << SP << SP << SP << "float a = " << fAttrActivationAlpha[direction * 3] << " * " << OpName
1108 << "_forget_gate[i] + " << fAttrActivationBeta[direction * 3] << ";\n";
1109 out << SP << SP << SP << "float b = (a > 0.) ? a : 0.;\n";
1110 }
1111 out << SP << SP << SP << SP << OpName << "_forget_gate[i] = (b < 1.) ? b : 1.;\n";
1112 out << SP << SP << "}\n";
1113 } else if (fAttrActivations[direction * 3] == "LeakyRelu") {
1114 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1115 out << SP << SP << SP << "if (" << OpName << "_forget_gate[i] < 0.)\n";
1116 out << SP << SP << SP << SP << OpName << "_forget_gate[i] = " << fAttrActivationAlpha[direction * 3]
1117 << " * " << OpName << "_forget_gate[i];\n";
1118 out << SP << SP << "}\n";
1119 } else if (fAttrActivations[direction * 3] == "ThresholdRelu") {
1120 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1121 out << SP << SP << SP << "if (" << OpName << "_forget_gate[i] < " << fAttrActivationAlpha[direction * 3]
1122 << ")\n";
1123 out << SP << SP << SP << SP << OpName << "_forget_gate[i] = 0.;\n";
1124 out << SP << SP << "}";
1125 } else if (fAttrActivations[direction * 3] == "Elu") {
1126 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1127 out << SP << SP << SP << "if (" << OpName << "_forget_gate[i] < 0.)\n";
1128 out << SP << SP << SP << SP << OpName << "_forget_gate[i] = " << fAttrActivationAlpha[direction * 3]
1129 << " * exp(" << OpName << "_forget_gate[i] - 1.);\n";
1130 out << SP << SP << "}\n";
1131 } else if (fAttrActivations[direction * 3] == "Softsign") {
1132 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1133 out << SP << SP << SP << SP << OpName << "_forget_gate[i] = " << OpName << "_forget_gate[i] / (1. + abs("
1134 << OpName << "_forget_gate[i]));\n";
1135 out << SP << SP << "}\n";
1136 } else { // fAttrActivations[direction * 3] = Softplus
1137 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1138 out << SP << SP << SP << SP << OpName << "_forget_gate[i] = log(1. + exp(" << OpName
1139 << "_forget_gate[i]));\n";
1140 out << SP << SP << "}\n";
1141 }
1142 }
1143
1144 // cell_state = input_gate o cell_gate
1145 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1146 out << SP << SP << SP << OpName << "_cell_state[i] = " << OpName << "_input_gate[i] * " << OpName
1147 << "_cell_gate[i];\n";
1148 out << SP << SP << "}\n";
1149
1150 if (fAttrInputForget == 0) {
1151 out << SP << SP << "if (seq == 0) {\n";
1152 if (!fNInitial_c.empty()) {
1153 // cell_state += forget_gate o initial_cell_state
1154 out << SP << SP << SP << "for (size_t i = 0; i < " << size << "; i++) {\n";
1155 out << SP << SP << SP << SP << OpName << "_cell_state[i + offset] += " << OpName
1156 << "_forget_gate[i + offset] * " << OpName << "_initial_cell_state[i];\n";
1157 out << SP << SP << SP << "}\n";
1158 }
1159 out << SP << SP << "} else {\n";
1160 // cell_state += forget_gate o previous_cell_state
1161 if (direction == 0) {
1162 if (fAttrDirection == "backward") {
1163 out << SP << SP << SP << "size_t previous_offset = (index + 1) * "
1164 << num_directions * batch_size * fAttrHiddenSize << ";\n";
1165 } else {
1166 out << SP << SP << SP << "size_t previous_offset = (seq - 1) * "
1167 << num_directions * batch_size * fAttrHiddenSize << ";\n";
1168 }
1169 } else { // direction=1
1170 out << SP << SP << SP << "size_t previous_offset = (index + 1) * "
1171 << num_directions * batch_size * fAttrHiddenSize << " + " << batch_size * fAttrHiddenSize << ";\n";
1172 }
1173 out << SP << SP << SP << "for (size_t i = 0; i < " << size << "; i++) {\n";
1174 out << SP << SP << SP << SP << OpName << "_cell_state[i + offset] += " << OpName
1175 << "_forget_gate[i + offset] * " << OpName << "_cell_state[i + previous_offset];\n";
1176 out << SP << SP << SP << "}\n";
1177 out << SP << SP << "}\n";
1178 }
1179
1180 if (!fNP.empty()) {
1181 // Peephole connection for the output gate
1182 if (direction == 0) {
1183 size_t p_offset = 2 * batch_size * fAttrHiddenSize;
1184 out << SP << SP << SP << "for (size_t i = 0; i < " << size << "; i++) {\n";
1185 out << SP << SP << SP << SP << OpName << "_output_gate[i + offset] += tensor_" << fNP << "[i + " << p_offset
1186 << "] * " << OpName << "_cell_state[i + offset];\n";
1187 out << SP << SP << SP << "}\n";
1188 } else { // direction=1
1189 size_t p_offset = 3 * batch_size * fAttrHiddenSize + 2 * batch_size * fAttrHiddenSize;
1190 out << SP << SP << SP << "for (size_t i = 0; i < " << size << "; i++) {\n";
1191 out << SP << SP << SP << SP << OpName << "_output_gate[i + offset] += tensor_" << fNP << "[i + " << p_offset
1192 << "] * " << OpName << "_cell_state[i + offset];\n";
1193 out << SP << SP << SP << "}\n";
1194 }
1195 }
1196
1197 // Clip the elements of the output gate into the range [-fAttrClip, fAttrClip]
1198 if (fAttrClip > .0) {
1199 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1200 if (fType == "float") {
1201 out << SP << SP << SP << "float x = (" << OpName << "_output_gate[i] > " << -fAttrClip << ") ? " << OpName
1202 << "_output_gate[i] : " << -fAttrClip << ";\n";
1203 }
1204 out << SP << SP << SP << OpName << "_output_gate[i] = (x < " << fAttrClip << ") ? x : " << fAttrClip << ";\n";
1205 out << SP << SP << "}\n";
1206 }
1207 // Apply the activation function to the output gate
1208 if (fAttrActivations[direction * 3] == "Relu") {
1209 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1210 out << SP << SP << SP << "if (" << OpName << "_output_gate[i] < 0.)\n";
1211 out << SP << SP << SP << SP << OpName << "_output_gate[i] = 0.;\n";
1212 out << SP << SP << "}\n";
1213 } else if (fAttrActivations[direction * 3] == "Tanh") {
1214 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1215 if (fType == "float") {
1216 out << SP << SP << SP << "float ex = exp(-2 * " << OpName << "_output_gate[i]);\n";
1217 }
1218 out << SP << SP << SP << SP << OpName << "_output_gate[i] = (1. - ex) / (1. + ex);\n";
1219 out << SP << SP << "}\n";
1220 } else if (fAttrActivations[direction * 3] == "Sigmoid") {
1221 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1222 out << SP << SP << SP << SP << OpName << "_output_gate[i] = 1. / (1. + exp(-" << OpName
1223 << "_output_gate[i]));\n";
1224 out << SP << SP << "}\n";
1225 } else if (fAttrActivations[direction * 3] == "Affine") {
1226 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1227 out << SP << SP << SP << SP << OpName << "_output_gate[i] = " << fAttrActivationAlpha[direction * 3] << " * "
1228 << OpName << "_output_gate[i] + " << fAttrActivationBeta[direction * 3] << ";\n";
1229 out << SP << SP << "}\n";
1230 } else if (fAttrActivations[direction * 3] == "ScaledTanh") {
1231 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1232 if (fType == "float") {
1233 out << SP << SP << SP << "float ex = exp(-2 * " << fAttrActivationBeta[direction * 3] << " * " << OpName
1234 << "_output_gate[i]);\n";
1235 }
1236 out << SP << SP << SP << SP << OpName << "_output_gate[i] = " << fAttrActivationAlpha[direction * 3]
1237 << " * (1. - ex) / (1. + ex);\n";
1238 out << SP << SP << "}\n";
1239 } else if (fAttrActivations[direction * 3] == "HardSigmoid") {
1240 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1241 if (fType == "float") {
1242 out << SP << SP << SP << "float a = " << fAttrActivationAlpha[direction * 3] << " * " << OpName
1243 << "_output_gate[i] + " << fAttrActivationBeta[direction * 3] << ";\n";
1244 out << SP << SP << SP << "float b = (a > 0.) ? a : 0.;\n";
1245 }
1246 out << SP << SP << SP << SP << OpName << "_output_gate[i] = (b < 1.) ? b : 1.;\n";
1247 out << SP << SP << "}\n";
1248 } else if (fAttrActivations[direction * 3] == "LeakyRelu") {
1249 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1250 out << SP << SP << SP << "if (" << OpName << "_output_gate[i] < 0.)\n";
1251 out << SP << SP << SP << SP << OpName << "_output_gate[i] = " << fAttrActivationAlpha[direction * 3] << " * "
1252 << OpName << "_output_gate[i];\n";
1253 out << SP << SP << "}\n";
1254 } else if (fAttrActivations[direction * 3] == "ThresholdRelu") {
1255 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1256 out << SP << SP << SP << "if (" << OpName << "_output_gate[i] < " << fAttrActivationAlpha[direction * 3]
1257 << ")\n";
1258 out << SP << SP << SP << SP << OpName << "_output_gate[i] = 0.;\n";
1259 out << SP << SP << "}";
1260 } else if (fAttrActivations[direction * 3] == "Elu") {
1261 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1262 out << SP << SP << SP << "if (" << OpName << "_output_gate[i] < 0.)\n";
1263 out << SP << SP << SP << SP << OpName << "_output_gate[i] = " << fAttrActivationAlpha[direction * 3]
1264 << " * exp(" << OpName << "_output_gate[i] - 1.);\n";
1265 out << SP << SP << "}\n";
1266 } else if (fAttrActivations[direction * 3] == "Softsign") {
1267 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1268 out << SP << SP << SP << SP << OpName << "_output_gate[i] = " << OpName << "_output_gate[i] / (1. + abs("
1269 << OpName << "_output_gate[i]));\n";
1270 out << SP << SP << "}\n";
1271 } else { // fAttrActivations[direction * 3] = Softplus
1272 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1273 out << SP << SP << SP << SP << OpName << "_output_gate[i] = log(1. + exp(" << OpName << "_output_gate[i]));\n";
1274 out << SP << SP << "}\n";
1275 }
1276
1277 // copy cell_state into new_cell_state
1278 out << SP << SP << "std::copy(" << OpName << "_cell_state + offset, " << OpName << "_cell_state + offset + "
1279 << size << ", " << OpName << "_new_cell_state + offset);\n";
1280 // Clip the elements of the new_cell_state into the range [-fAttrClip, fAttrClip]
1281 if (fAttrClip > .0) {
1282 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1283 if (fType == "float") {
1284 out << SP << SP << SP << "float x = (" << OpName << "_new_cell_state[i] > " << -fAttrClip << ") ? "
1285 << OpName << "_new_cell_state[i] : " << -fAttrClip << ";\n";
1286 }
1287 out << SP << SP << SP << OpName << "_new_cell_state[i] = (x < " << fAttrClip << ") ? x : " << fAttrClip
1288 << ";\n";
1289 out << SP << SP << "}\n";
1290 }
1291 // Apply the activation function to the new cell state
1292 if (fAttrActivations[direction * 3 + 2] == "Relu") {
1293 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1294 out << SP << SP << SP << "if (" << OpName << "_new_cell_state[i] < 0.)\n";
1295 out << SP << SP << SP << SP << OpName << "_new_cell_state[i] = 0.;\n";
1296 out << SP << SP << "}\n";
1297 } else if (fAttrActivations[direction * 3 + 2] == "Tanh") {
1298 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1299 if (fType == "float") {
1300 out << SP << SP << SP << "float ex = exp(-2 * " << OpName << "_new_cell_state[i]);\n";
1301 }
1302 out << SP << SP << SP << SP << OpName << "_new_cell_state[i] = (1. - ex) / (1. + ex);\n";
1303 out << SP << SP << "}\n";
1304 } else if (fAttrActivations[direction * 3 + 2] == "Sigmoid") {
1305 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1306 out << SP << SP << SP << SP << OpName << "_new_cell_state[i] = 1. / (1. + exp(-" << OpName
1307 << "_new_cell_state[i]));\n";
1308 out << SP << SP << "}\n";
1309 } else if (fAttrActivations[direction * 3 + 2] == "Affine") {
1310 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1311 out << SP << SP << SP << SP << OpName << "_new_cell_state[i] = " << fAttrActivationAlpha[direction * 3 + 2]
1312 << " * " << OpName << "_new_cell_state[i] + " << fAttrActivationBeta[direction * 3 + 2] << ";\n";
1313 out << SP << SP << "}\n";
1314 } else if (fAttrActivations[direction * 3 + 2] == "ScaledTanh") {
1315 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1316 if (fType == "float") {
1317 out << SP << SP << SP << "float ex = exp(-2 * " << fAttrActivationBeta[direction * 3 + 2] << " * " << OpName
1318 << "_new_cell_state[i]);\n";
1319 }
1320 out << SP << SP << SP << SP << OpName << "_new_cell_state[i] = " << fAttrActivationAlpha[direction * 3 + 2]
1321 << " * (1. - ex) / (1. + ex);\n";
1322 out << SP << SP << "}\n";
1323 } else if (fAttrActivations[direction * 3 + 2] == "HardSigmoid") {
1324 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1325 if (fType == "float") {
1326 out << SP << SP << SP << "float a = " << fAttrActivationAlpha[direction * 3 + 2] << " * " << OpName
1327 << "_new_cell_state[i] + " << fAttrActivationBeta[direction * 3 + 2] << ";\n";
1328 out << SP << SP << SP << "float b = (a > 0.) ? a : 0.;\n";
1329 }
1330 out << SP << SP << SP << SP << OpName << "_new_cell_state[i] = (b < 1.) ? b : 1.;\n";
1331 out << SP << SP << "}\n";
1332 } else if (fAttrActivations[direction * 3 + 2] == "LeakyRelu") {
1333 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1334 out << SP << SP << SP << "if (" << OpName << "_new_cell_state[i] < 0.)\n";
1335 out << SP << SP << SP << SP << OpName << "_new_cell_state[i] = " << fAttrActivationAlpha[direction * 3 + 2]
1336 << " * " << OpName << "_new_cell_state[i];\n";
1337 out << SP << SP << "}\n";
1338 } else if (fAttrActivations[direction * 3 + 2] == "ThresholdRelu") {
1339 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1340 out << SP << SP << SP << "if (" << OpName << "_new_cell_state[i] < " << fAttrActivationAlpha[direction * 3 + 2]
1341 << ")\n";
1342 out << SP << SP << SP << SP << OpName << "_new_cell_state[i] = 0.;\n";
1343 out << SP << SP << "}";
1344 } else if (fAttrActivations[direction * 3 + 2] == "Elu") {
1345 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1346 out << SP << SP << SP << "if (" << OpName << "_new_cell_state[i] < 0.)\n";
1347 out << SP << SP << SP << SP << OpName << "_new_cell_state[i] = " << fAttrActivationAlpha[direction * 3 + 2]
1348 << " * exp(" << OpName << "_new_cell_state[i] - 1.);\n";
1349 out << SP << SP << "}\n";
1350 } else if (fAttrActivations[direction * 3 + 2] == "Softsign") {
1351 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1352 out << SP << SP << SP << SP << OpName << "_new_cell_state[i] = " << OpName << "_new_cell_state[i] / (1. + abs("
1353 << OpName << "_new_cell_state[i]));\n";
1354 out << SP << SP << "}\n";
1355 } else { // fAttrActivations[direction * 3 + 2] = Softplus
1356 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1357 out << SP << SP << SP << SP << OpName << "_new_cell_state[i] = log(1. + exp(" << OpName
1358 << "_new_cell_state[i]));\n";
1359 out << SP << SP << "}\n";
1360 }
1361
1362 // hidden_state = output_gate o new_cell_state
1363 out << SP << SP << "for (size_t i = offset; i < offset + " << size << "; i++) {\n";
1364 out << SP << SP << SP << OpName << "_hidden_state[i] = " << OpName << "_output_gate[i] * " << OpName
1365 << "_new_cell_state[i];\n";
1366 out << SP << SP << "}\n";
1367 out << SP << "}\n";
1368 }
1369
1370 // Padding the hidden state for LSTM with different sequence lengths
1371 if (!fNSequence_lens.empty()) {
1372 out << SP << "for (size_t seq = 0; seq < " << seq_length << "; seq++) {\n";
1373 out << SP << SP << "for (size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
1374 out << SP << SP << SP << "if (seq >= tensor_" << fNSequence_lens << "[batch]) {\n";
1375 for (size_t direction = 0; direction < num_directions; direction++) {
1376 out << SP << SP << SP << SP << SP << "for (size_t h = 0; h < " << fAttrHiddenSize << "; h++) {\n";
1377 out << SP << SP << SP << SP << SP << SP << "size_t idx = seq * "
1378 << num_directions * batch_size * fAttrHiddenSize + direction * batch_size * fAttrHiddenSize
1379 << " + batch * " << fAttrHiddenSize << " + h;\n";
1380 out << SP << SP << SP << SP << SP << SP << OpName << "_cell_state[idx] = 0.;\n";
1381 out << SP << SP << SP << SP << SP << SP << OpName << "_hidden_state[idx] = 0.;\n";
1382 out << SP << SP << SP << SP << SP << "}\n";
1383 }
1384 out << SP << SP << SP << "}\n";
1385 out << SP << SP << "}\n";
1386 out << SP << "}\n";
1387 }
1388
1389 // Copy the hidden state into y and y_h and copy cell_state into y_c
1390 if (fAttrLayout == 0) {
1391 if (!fNY_h.empty()) {
1392 // Copy hidden_state into Y_h
1393 if (fNSequence_lens.empty()) {
1394 size_t y_h_size = batch_size * fAttrHiddenSize;
1395 if (fAttrDirection == "backward") {
1396 out << SP << "std::copy(" << OpName << "_hidden_state, " << OpName << "_hidden_state + " << y_h_size
1397 << ", tensor_" << fNY_h << ");\n";
1398 } else {
1399 size_t offset = (seq_length - 1) * num_directions * batch_size * fAttrHiddenSize;
1400 out << SP << "std::copy(" << OpName << "_hidden_state + " << offset << ", " << OpName
1401 << "_hidden_state + " << offset << " + " << y_h_size << ", tensor_" << fNY_h << ");\n";
1402 }
1403 if (num_directions == 2) {
1404 out << SP << "std::copy(" << OpName << "_hidden_state + " << y_h_size << ", " << OpName
1405 << "_hidden_state + " << 2 * y_h_size << ", tensor_" << fNY_h << " + " << y_h_size << ");\n";
1406 }
1407 } else { // LSTM with different sequence lengths
1408 if (fAttrDirection == "backward") {
1409 out << SP << "for (size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
1410 out << SP << SP << "size_t offset = batch * " << fAttrHiddenSize << ";\n";
1411 out << SP << SP << "std::copy(" << OpName << "_hidden_state + offset, " << OpName
1412 << "_hidden_state + offset + " << fAttrHiddenSize << ", tensor_" << fNY_h << " + offset);\n";
1413 out << SP << "}\n";
1414 } else {
1415 out << SP << "for (size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
1416 out << SP << SP << "size_t seq = " << "tensor_" << fNSequence_lens << "[batch] - 1;\n";
1417 out << SP << SP << "size_t offset = seq * " << num_directions * batch_size * fAttrHiddenSize
1418 << " + batch * " << fAttrHiddenSize << ";\n";
1419 out << SP << SP << "size_t y_h_offset = batch * " << fAttrHiddenSize << ";\n";
1420 out << SP << SP << "std::copy(" << OpName << "_hidden_state + offset, " << OpName
1421 << "_hidden_state + offset + " << fAttrHiddenSize << ", tensor_" << fNY_h << " + y_h_offset);\n";
1422 out << SP << "}\n";
1423 }
1424 if (num_directions == 2) {
1425 out << SP << "for (size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
1426 out << SP << SP << "size_t offset = " << batch_size * fAttrHiddenSize << " + batch * " << fAttrHiddenSize
1427 << ";\n";
1428 out << SP << SP << "size_t y_h_offset = " << batch_size * fAttrHiddenSize << " + batch * "
1429 << fAttrHiddenSize << ";\n";
1430 out << SP << SP << "std::copy(" << OpName << "_hidden_state + offset, " << OpName
1431 << "_hidden_state + offset + " << fAttrHiddenSize << ", tensor_" << fNY_h << " + y_h_offset);\n";
1432 out << SP << "}\n";
1433 }
1434 }
1435 }
1436 if (!fNY_c.empty()) {
1437 // Copy cell_state into Y_c
1438 if (fNSequence_lens.empty()) {
1439 size_t y_h_size = batch_size * fAttrHiddenSize;
1440 if (fAttrDirection == "backward") {
1441 out << SP << "std::copy(" << OpName << "_cell_state, " << OpName << "_hidden_state + " << y_h_size
1442 << ", tensor_" << fNY_c << ");\n";
1443 } else {
1444 size_t offset = (seq_length - 1) * num_directions * batch_size * fAttrHiddenSize;
1445 out << SP << "std::copy(" << OpName << "_cell_state + " << offset << ", " << OpName << "_cell_state + "
1446 << offset << " + " << y_h_size << ", tensor_" << fNY_c << ");\n";
1447 }
1448 if (num_directions == 2) {
1449 out << SP << "std::copy(" << OpName << "_cell_state + " << y_h_size << ", " << OpName << "_cell_state + "
1450 << 2 * y_h_size << ", tensor_" << fNY_c << " + " << y_h_size << ");\n";
1451 }
1452 } else { // LSTM with different sequence lengths
1453 if (fAttrDirection == "backward") {
1454 out << SP << "for (size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
1455 out << SP << SP << "size_t offset = batch * " << fAttrHiddenSize << ";\n";
1456 out << SP << SP << "std::copy(" << OpName << "_cell_state + offset, " << OpName
1457 << "_cell_state + offset + " << fAttrHiddenSize << ", tensor_" << fNY_c << " + offset);\n";
1458 out << SP << "}\n";
1459 } else {
1460 out << SP << "for (size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
1461 out << SP << SP << "size_t seq = " << "tensor_" << fNSequence_lens << "[batch] - 1;\n";
1462 out << SP << SP << "size_t offset = seq * " << num_directions * batch_size * fAttrHiddenSize
1463 << " + batch * " << fAttrHiddenSize << ";\n";
1464 out << SP << SP << "size_t y_h_offset = batch * " << fAttrHiddenSize << ";\n";
1465 out << SP << SP << "std::copy(" << OpName << "_cell_state + offset, " << OpName
1466 << "_cell_state + offset + " << fAttrHiddenSize << ", tensor_" << fNY_c << " + y_h_offset);\n";
1467 out << SP << "}\n";
1468 }
1469 if (num_directions == 2) {
1470 out << SP << "for (size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
1471 out << SP << SP << "size_t offset = " << batch_size * fAttrHiddenSize << " + batch * " << fAttrHiddenSize
1472 << ";\n";
1473 out << SP << SP << "size_t y_h_offset = " << batch_size * fAttrHiddenSize << " + batch * "
1474 << fAttrHiddenSize << ";\n";
1475 out << SP << SP << "std::copy(" << OpName << "_cell_state + offset, " << OpName
1476 << "_cell_state + offset + " << fAttrHiddenSize << ", tensor_" << fNY_c << " + y_h_offset);\n";
1477 out << SP << "}\n";
1478 }
1479 }
1480 }
1481 } else { // fAttrLayout=1
1482 if (!fNY.empty()) {
1483 // Copy hidden_state into Y
1484 for (size_t direction = 0; direction < num_directions; direction++) {
1485 out << SP << "for (size_t seq = 0; seq < " << seq_length << "; seq++) {\n";
1486 out << SP << SP << "for (size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
1487 out << SP << SP << SP << "size_t offset = seq * " << num_directions * batch_size * fAttrHiddenSize << " + "
1488 << direction * batch_size * fAttrHiddenSize << " + batch * " << fAttrHiddenSize << ";\n";
1489 out << SP << SP << SP << "size_t y_offset = batch * " << seq_length * num_directions * fAttrHiddenSize
1490 << " + seq * " << num_directions * fAttrHiddenSize << " + " << direction * fAttrHiddenSize << ";\n";
1491 out << SP << SP << SP << "std::copy(" << OpName << "_hidden_state + offset, " << OpName
1492 << "_hidden_state + offset + " << fAttrHiddenSize << ", tensor_" << fNY << " + y_offset);\n";
1493 out << SP << SP << "}\n";
1494 out << SP << "}\n";
1495 }
1496 }
1497 if (!fNY_h.empty()) {
1498 // Copy the hidden_state into Y_h
1499 if (fAttrDirection == "backward") {
1500 out << SP << "for (size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
1501 out << SP << SP << "size_t offset = batch * " << fAttrHiddenSize << ";\n";
1502 out << SP << SP << "size_t y_h_offset = batch * " << num_directions * fAttrHiddenSize << ";\n";
1503 out << SP << SP << "std::copy(" << OpName << "_hidden_state + offset, " << OpName
1504 << "_hidden_state + offset + " << fAttrHiddenSize << ", tensor_" << fNY_h << " + y_h_offset);\n";
1505 out << SP << "}\n";
1506 } else {
1507 out << SP << "for (size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
1508 if (fNSequence_lens.empty()) {
1509 out << SP << SP << "size_t seq = " << seq_length - 1 << ";\n";
1510 } else {
1511 out << SP << SP << "size_t seq = " << "tensor_" << fNSequence_lens << "[batch] - 1;\n";
1512 }
1513 out << SP << SP << "size_t offset = seq * " << num_directions * batch_size * fAttrHiddenSize
1514 << " + batch * " << fAttrHiddenSize << ";\n";
1515 out << SP << SP << "size_t y_h_offset = batch * " << num_directions * fAttrHiddenSize << ";\n";
1516 out << SP << SP << "std::copy(" << OpName << "_hidden_state + offset, " << OpName
1517 << "_hidden_state + offset + " << fAttrHiddenSize << ", tensor_" << fNY_h << " + y_h_offset);\n";
1518 out << SP << "}\n";
1519 }
1520 if (num_directions == 2) {
1521 out << SP << "for (size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
1522 out << SP << SP << "size_t offset = " << batch_size * fAttrHiddenSize << " + batch * " << fAttrHiddenSize
1523 << ";\n";
1524 out << SP << SP << "size_t y_h_offset = batch * " << num_directions * fAttrHiddenSize << " + "
1525 << fAttrHiddenSize << ";\n";
1526 out << SP << SP << "std::copy(" << OpName << "_hidden_state + offset, " << OpName
1527 << "_hidden_state + offset + " << fAttrHiddenSize << ", tensor_" << fNY_h << " + y_h_offset);\n";
1528 out << SP << "}\n";
1529 }
1530 }
1531
1532 if (!fNY_c.empty()) {
1533 // copy the cell_state into Y_c
1534 if (fAttrDirection == "backward") {
1535 out << SP << "for (size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
1536 out << SP << SP << "size_t offset = batch * " << fAttrHiddenSize << ";\n";
1537 out << SP << SP << "size_t y_h_offset = batch * " << num_directions * fAttrHiddenSize << ";\n";
1538 out << SP << SP << "std::copy(" << OpName << "_cell_state + offset, " << OpName << "_cell_state + offset + "
1539 << fAttrHiddenSize << ", tensor_" << fNY_c << " + y_h_offset);\n";
1540 out << SP << "}\n";
1541 } else {
1542 out << SP << "for (size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
1543 if (fNSequence_lens.empty()) {
1544 out << SP << SP << "size_t seq = " << seq_length - 1 << ";\n";
1545 } else {
1546 out << SP << SP << "size_t seq = " << "tensor_" << fNSequence_lens << "[batch] - 1;\n";
1547 }
1548 out << SP << SP << "size_t offset = seq * " << num_directions * batch_size * fAttrHiddenSize
1549 << " + batch * " << fAttrHiddenSize << ";\n";
1550 out << SP << SP << "size_t y_h_offset = batch * " << num_directions * fAttrHiddenSize << ";\n";
1551 out << SP << SP << "std::copy(" << OpName << "_cell_state + offset, " << OpName << "_cell_state + offset + "
1552 << fAttrHiddenSize << ", tensor_" << fNY_c << " + y_h_offset);\n";
1553 out << SP << "}\n";
1554 }
1555 if (num_directions == 2) {
1556 out << SP << "for (size_t batch = 0; batch < " << batch_size << "; batch++) {\n";
1557 out << SP << SP << "size_t offset = " << batch_size * fAttrHiddenSize << " + batch * " << fAttrHiddenSize
1558 << ";\n";
1559 out << SP << SP << "size_t y_h_offset = batch * " << num_directions * fAttrHiddenSize << " + "
1560 << fAttrHiddenSize << ";\n";
1561 out << SP << SP << "std::copy(" << OpName << "_cell_state + offset, " << OpName << "_cell_state + offset + "
1562 << fAttrHiddenSize << ", tensor_" << fNY_c << " + y_h_offset);\n";
1563 out << SP << "}\n";
1564 }
1565 }
1566 }
1567
1568 return out.str();
1569}
1570
1571} // namespace TMVA::Experimental::SOFIE
1572
1573#endif
#define b(i)
Definition RSha256.hxx:100
#define h(i)
Definition RSha256.hxx:106
size_t size(const MatrixT &matrix)
retrieve the size of a square matrix
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void input
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void char Point_t Rectangle_t WindowAttributes_t Float_t Float_t Float_t Int_t Int_t UInt_t UInt_t Rectangle_t Int_t Int_t Window_t TString Int_t GCValues_t GetPrimarySelectionOwner GetDisplay GetScreen GetColormap GetNativeEvent const char const char dpyName wid window const char font_name cursor keysym reg const char only_if_exist regb h Point_t winding char text const char depth char const char Int_t count const char ColorStruct_t color const char Pixmap_t Pixmap_t PictureAttributes_t attr const char char ret_data h unsigned char height h offset
char name[80]
Definition TGX11.cxx:142
const_iterator begin() const
const_iterator end() const
Long Short-Term Memory operator.
std::string GenerateSessionMembersCode(std::string opName) override
Generate the code for the Session internal data vectors.
std::vector< size_t > fShapeR
Shape of the recurrence.
std::vector< size_t > fShapeInitial_c
Shape of the initial value of the cell states.
std::string fNY_c
Name of the last sequence of the cell states.
ROperator_LSTM(std::vector< float > activation_alpha, std::vector< float > activation_beta, std::vector< std::string > activations, float clip, std::string direction, size_t hidden_size, size_t input_forget, size_t layout, std::string nameX, std::string nameW, std::string nameR, std::string nameB, std::string nameSequence_lens, std::string nameInitial_h, std::string nameInitial_c, std::string nameP, std::string nameY, std::string nameY_h, std::string nameY_c)
Constructor of ROperator_LSTM from the attributes.
std::string fNR
Name of the recurrence.
std::vector< size_t > fShapeY_h
Shape of the last sequence of the output.
std::vector< float > fAttrActivationAlpha
Sacling values used by some activation functions.
std::vector< size_t > fShapeInitial_h
Shape of the initial value of the hidden states.
std::vector< std::vector< size_t > > ShapeInference(std::vector< std::vector< size_t > > input) override
Infers the shape of the output tensors.
size_t fAttrHiddenSize
Number of the hidden layers.
std::vector< std::string > GetBlasRoutines() override
Returns the blas routines needed to compile the generated code.
std::string fNInitial_c
Name of the initial value of the cell states.
std::vector< size_t > fShapeB
Shape of the bias.
std::string fNW
Name of the weights.
std::vector< size_t > fShapeY
Shape of the output.
std::vector< size_t > fShapeP
Shape of the peepholes.
std::string fType
Type of the tensors.
std::string fAttrDirection
Direction of processing.
std::vector< float > fAttrActivationBeta
Scaling values used by some activation functions.
std::vector< size_t > fShapeX
Shape of the input.
std::vector< size_t > fShapeW
Shape of the weights.
std::vector< ETensorType > TypeInference(std::vector< ETensorType > input) override
Infers the type of the output tensors.
std::string fNSequence_lens
Name of length of the sequences.
std::string fNY_h
Name of the last sequence of the output.
std::string fNY
Name of the output.
void Initialize(RModel &) override
Initialize the model.
std::vector< std::string > fAttrActivations
Activation functions.
std::string Generate(std::string OpName) override
Generate the inference code.
std::vector< size_t > fShapeSequence_lens
Shape of the length of the sequences.
std::vector< size_t > fShapeY_c
Shape of the last sequence of the cell states.
std::string fNInitial_h
Name of the initial value of the hidden states.
ROperator_LSTM()
Default constructor of ROperator_LSTM.
std::vector< std::string_view > fInputTensorNames
Definition ROperator.hxx:46
std::vector< std::string_view > fOutputTensorNames
Definition ROperator.hxx:47
static uint64_t sum(uint64_t i)
Definition Factory.cxx:2335