Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
ROperator_GRU.hxx
Go to the documentation of this file.
1#ifndef TMVA_SOFIE_ROPERATOR_GRU
2#define TMVA_SOFIE_ROPERATOR_GRU
3
4#include "TMVA/RModel.hxx"
5#include "TMVA/ROperator.hxx"
7
8#include <memory>
9#include <sstream>
10#include <stdexcept>
11#include <string>
12#include <vector>
13
15
16/*! \brief Gated Recurrent Unit operator
17 *
18 * Inference code generation for one-layer GRU. Supports forward, reverse and bidirectional GRU.
19 * See the <a href="https://github.com/onnx/onnx/blob/master/docs/Operators.md#GRU">ONNX documentation</a>
20 * for details about the supported GRU architectures.
21 */
22template <typename T> class ROperator_GRU final : public ROperator {
23 private:
24 std::vector<float> fAttrActivationAlpha; ///< Scaling values used by some activation functions
25 std::vector<float> fAttrActivationBeta; ///< Scaling values used by some activation functions
26 std::vector<std::string> fAttrActivations; ///< Activation functions
27 float fAttrClip; ///< Clip threshold
28 std::string fAttrDirection; ///< Direction of processing
29 size_t fAttrHiddenSize; ///< Number of the hidden layers
30 size_t fAttrLayout; ///< Data layout
31 size_t fAttrLinearBeforeReset; ///< Linear layer before the reset gate
32
33 std::string fNX; ///< Name of the input
34 std::string fNW; ///< Name of the weights
35 std::string fNR; ///< Name of the recurrence
36 std::string fNB; ///< Name of the bias
37 std::string fNSequence_lens; ///< Name of the length of the sequences
38 std::string fNInitial_h; ///< Name of the initial value of the hidden states
39 std::string fNY; ///< Name of the output
40 std::string fNY_h; ///< Name of the last sequence of the output
41
42 std::vector<size_t> fShapeX; ///< Shape of the input
43 std::vector<size_t> fShapeW; ///< Shape of the weights
44 std::vector<size_t> fShapeR; ///< Shape of the recurrence
45 std::vector<size_t> fShapeB; ///< Shape of the bias
46 std::vector<size_t> fShapeSequence_lens; ///< Shape of the length of the sequences
47 std::vector<size_t> fShapeInitial_h; ///< Shape of the initial value of hidden states
48 std::vector<size_t> fShapeY; ///< Shape of the output
49 std::vector<size_t> fShapeY_h; ///< Shape of the last sequence of the output
50
51 std::string fType; ///< Type of the tensors
52
53 public:
54 /*! Default constructor of ROperator_GRU */
56
57 /*! \brief Constructor of ROperator_GRU from the attributes
58 *
59 * \param activation_alpha scaling values used by some activation functions
60 * \param activation_beta scaling values used by some activation functions
61 * \param activations activation functions
62 * \param clip clip threshold
63 * \param direction direction of processing of the sequneces
64 * \param hidden_size number of hidden layers
65 * \param layout data layout
66 * \param linear_before_reset Linear layer before the reset gate
67 * \param nameX name of the input tensor
68 * \param nameW name of the weight tensor
69 * \param nameR name of the recurrence tensor
70 * \param nameB name of the bias tensor
71 * \param nameSequence_lens name of the length of the sequences
72 * \param nameInitial_h name of the initial value of the hidden states
73 * \param nameY name of the output
74 * \param nameY_h name of the last sequence of the output
75 */
76 ROperator_GRU(std::vector<float> activation_alpha,
77 std::vector<float> activation_beta,
78 std::vector<std::string> activations, float clip,
79 std::string direction, size_t hidden_size,
80 size_t layout, size_t linear_before_reset,
81 std::string nameX, std::string nameW, std::string nameR,
82 std::string nameB, std::string nameSequence_lens,
83 std::string nameInitial_h, std::string nameY, std::string nameY_h)
88 fNX(UTILITY::Clean_name(nameX)), fNW(UTILITY::Clean_name(nameW)),
89 fNR(UTILITY::Clean_name(nameR)), fNB(UTILITY::Clean_name(nameB)),
90 fNSequence_lens(UTILITY::Clean_name(nameSequence_lens)),
91 fNInitial_h(UTILITY::Clean_name(nameInitial_h)),
92 fNY(UTILITY::Clean_name(nameY)), fNY_h(UTILITY::Clean_name(nameY_h)) {
93
95 if (!fNB.empty()){
96 fInputTensorNames.emplace_back(fNB);
97 }
98 if (!fNSequence_lens.empty()){
100 }
101 if (!fNInitial_h.empty()){
102 fInputTensorNames.emplace_back(fNInitial_h);
103 }
104
105 fOutputTensorNames = { };
106 if (!fNY.empty()){
107 fOutputTensorNames.emplace_back(fNY);
108 }
109 if (!fNY_h.empty()){
110 fOutputTensorNames.emplace_back(fNY_h);
111 }
112
113 if (std::is_same<T, float>::value) {
114 fType = "float";
115 } else {
116 throw std::runtime_error(
117 "TMVA SOFIE Encountered unsupported type parsing a GRU operator");
118 }
119 }
120
121 /*! \brief Infers the shape of the output tensors
122 *
123 * \param input shape of the input tensors
124 */
125 std::vector<std::vector<size_t>> ShapeInference(std::vector<std::vector<size_t>> /*input*/);
126
127 /*! \brief Initialize the model
128 *
129 * \param model Model
130 */
131 void Initialize(RModel &) override;
132
133 /*! \brief Generate the inference code
134 *
135 * \param OpName name of the operator
136 */
137 std::string Generate(std::string /*OpName*/) override;
138
139 /*! \brief Returns the blas routines needed to compile the generated code
140 */
141 std::vector<std::string> GetBlasRoutines() override { return { std::string("Gemm"), std::string("Axpy") }; }
142};
143
144template <typename T>
145auto ROperator_GRU<T>::ShapeInference(std::vector<std::vector<size_t>> input) -> std::vector<std::vector<size_t>>
146{
147 size_t num_directions = input[1][0];
148 size_t hidden_size = input[1][1] / 3;
149 if (fAttrLayout == 0) {
150 size_t seq_length = input[0][0];
151 size_t batch_size = input[0][1];
152 std::vector<std::vector<size_t>> ret(
154 return ret;
155 } else {
156 size_t batch_size = input[0][0];
157 size_t seq_length = input[0][1];
158 std::vector<std::vector<size_t>> ret(
160 return ret;
161 }
162}
163
164template <typename T>
166{
167
168 // Check the input and output tensors
169 if (!model.CheckIfTensorAlreadyExist(fNX)) {
170 throw std::runtime_error("TMVA SOFIE GRU Op input tensor " + fNX + " is not found in model.");
171 }
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.");
175 }
176 if (!model.CheckIfTensorAlreadyExist(fNW)) {
177 throw std::runtime_error("TMVA SOFIE GRU Op input tensor " + fNW + " is not found in model.");
178 }
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.");
182 }
183 if (!model.CheckIfTensorAlreadyExist(fNR)) {
184 throw std::runtime_error("TMVA SOFIE GRU Op input tensor " + fNR + " is not found in model.");
185 }
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.");
189 }
190 if (!fNB.empty()) {
191 if (!model.CheckIfTensorAlreadyExist(fNB)) {
192 throw std::runtime_error("TMVA SOFIE GRU op input tensor " + fNB + " is not found in model.");
193 }
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.");
197 }
198 if (fShapeB.size() == 2) {
199 // Broadcasting the bias
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,
216 new_bias + offset);
217 }
218 }
219 }
220 }
221
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);
226 }
227 }
228 }
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.");
232 }
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.");
236 }
237 }
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.");
241 }
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.");
245 }
246 }
247 if (!fNY.empty()) {
248 fShapeY = ShapeInference({fShapeX, fShapeW})[0];
249 if (!model.CheckIfTensorAlreadyExist(fNY)) {
250 model.AddIntermediateTensor(fNY, model.GetTensorType(fNX), fShapeY);
251 }
252 }
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);
257 }
258 }
259 // Check the attributes
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");
265 }
266 }
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);
272 }
273 if (3 * fAttrHiddenSize != fShapeW[1]) {
274 throw std::runtime_error("TMVA SOFIE - fAttrHiddenSize must be equal to " + std::to_string(fShapeW[1] / 3));
275 }
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)");
279 }
280 if (fAttrLinearBeforeReset > 1) {
281 throw std::runtime_error("TMVA SOFIE - fAttrInputForget = " + std::to_string(fAttrLinearBeforeReset) +
282 " must be 0 or 1.");
283 }
284 if (fAttrActivations.empty()) {
285 if (fAttrDirection == "bidirectional") {
286 fAttrActivations = {"Sigmoid", "Tanh", "Sigmoid", "Tanh"};
287 } else {
288 fAttrActivations = {"Sigmoid", "Tanh"};
289 }
290 }
291
292 // To get unique intermediate tensor names, we add the name of the input
293 // tensor. One might also consider using the index of the operator in the
294 // RMode, but this information is not available in the current scope.
295 std::string opName = "op_gru_" + fNX;
296
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];
301
302 auto declareVector = [&](std::string const &name, std::size_t n) {
303 std::string fullName = opName + "_" + name;
304 model.AddIntermediateTensor(fullName, ConvertStringToType(fType), std::vector<std::size_t>{n});
305 };
306
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);
311 }
312 // Set the feedforward
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);
317 // gate results
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);
322
323 // feedback
324 declareVector("feedback", batch_size * fAttrHiddenSize);
325
326 // hiddden state
327 if (fAttrLayout != 0 || fNY.empty()) {
328 declareVector("hidden_state", hs_size);
329 }
330}
331
332template <typename T>
333auto ROperator_GRU<T>::Generate(std::string OpName) -> std::string
334{
335 OpName = "op_" + OpName;
336 std::stringstream out;
337
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];
342
343 auto getVec = [&](std::string const &name) { return "tensor_op_gru_" + fNX + "_" + name; };
344
345 // set the input
346 if (fAttrLayout == 0) {
347 out << SP << fType << " const* " << OpName << "_input = tensor_" << fNX << ";\n";
348 } else {
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
355 << " + i];\n";
356 out << SP << SP << SP << "}\n";
357 out << SP << SP << "}\n";
358 out << SP << "}\n";
359 }
360
361 // Set the initial hidden state
362 if (!fNInitial_h.empty()) {
363 if (fAttrLayout == 0) {
364 out << SP << fType << " *" << OpName << "_initial_hidden_state = " << " tensor_" << fNInitial_h << ";\n";
365 } else {
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";
374 out << SP << "}\n";
375 }
376 }
377 }
378
379 // Set the feedforward
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";
383 // Set the gates
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";
387 // Set the hidden state
388 if (fAttrLayout == 0 && !fNY.empty()) {
389 out << SP << fType << " *" << OpName << "_hidden_state = tensor_" << fNY << ";\n";
390 } else {
391 out << SP << fType << " * " << OpName << "_hidden_state = " << getVec("hidden_state") << ";\n";
392 }
393
394 out << SP << fType << " * " << OpName << "_feedback = " << getVec("feedback") << ";\n";
395
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";
405 }
406 if (!fNB.empty()) {
407 out << SP << "int " << OpName << "_bias_size = " << seq_length * batch_size * fAttrHiddenSize << ";\n";
408 }
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";
412
413 for (size_t direction = 0; direction < num_directions; direction++) {
414 if (direction == 0) {
415 if (fType == "float") {
416 // f_update_gate = input * weight_z^T
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";
421 // f_reset_gate = input * weight_r^T
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";
427 // f_hidden_gate = input * weight_h^T
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";
433 }
434 } else {
435 if (fType == "float") {
436 // f_update_gate = input * weight_z^T
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";
442 // f_reset_gate = input * weight_r^T
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";
448 // f_hidden_gate = input * weight_h^T
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";
454 }
455 }
456
457 if (!fNB.empty()) {
458 if (direction == 0) {
459 if (fType == "float") {
460 // Add the bias of the weight to f_update_gate
461 out << SP << "BLAS::saxpy_(&" << OpName << "_bias_size, &" << OpName << "_alpha, tensor_" << fNB << ", &"
462 << OpName << "_incx, " << OpName << "_f_update_gate, &" << OpName << "_incy);\n";
463 // Add the bias of the recurrence to f_update_gate
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
467 << "_incy);\n";
468 // Add the bias of the weight to f_reset_gate
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
472 << "_incy);\n";
473 // Add the bias of the recurrence to f_reset_gate
474 // size_t rbr_offset = fAttrHiddenSize * fAttrHiddenSize + 3 * batch_size * fAttrHiddenSize;
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
478 << "_incy);\n";
479 // Add the bias of the weight to f_hidden_gate
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
483 << "_incy);\n";
484 if (fAttrLinearBeforeReset == 0) {
485 // Add the bias of the recurrence to f_hidden_gate
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
489 << "_incy);\n";
490 }
491 }
492 } else {
493 if (fType == "float") {
494 // Add the bias of the weight to f_update_gate
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
498 << "_incy);\n";
499 // Add the bias of the recurrence to f_update_gate
500 // size_t rbz_offset = 3 * fAttrHiddenSize * fAttrHiddenSize + 3 * batch_size * fAttrHiddenSize;
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
504 << "_incy);\n";
505 // Add the bias of the weight to f_reset_gate
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
509 << "_incy);\n";
510 // Add the bias of the recurrence to f_reset_gate
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
514 << "_incy);\n";
515 // Add the bias of the weight to f_hidden_gate
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
519 << "_incy);\n";
520 if (fAttrLinearBeforeReset == 0) {
521 // Add the bias of the recurrence to f_hidden_gate
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
525 << "_incy);\n";
526 }
527 }
528 }
529 }
530
531 // Copy the feedforward into the gates
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";
536 } else {
537 out << SP << SP << "size_t gate_offset = seq * " << num_directions * batch_size * fAttrHiddenSize << " + "
538 << batch_size * fAttrHiddenSize << ";\n";
539 }
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";
547 out << SP << "}\n";
548
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";
552 } else {
553 out << SP << SP << "size_t index = seq;\n";
554 }
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";
558 } else {
559 out << SP << SP << "size_t offset = index * " << num_directions * batch_size * fAttrHiddenSize << " + "
560 << batch_size * fAttrHiddenSize << ";\n";
561 }
562 size_t size = batch_size * fAttrHiddenSize;
563 // gate = gate + initial_hidden_state * Recurrence^T
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";
577 }
578 } else { // direction=1
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";
590 }
591 }
592 }
593 out << SP << SP << "} else {\n";
594 // gate = gate + previous_hidden_state * Recurrence^T
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";
599 } else {
600 out << SP << SP << SP << "size_t previous_offset = (seq - 1) * "
601 << num_directions * batch_size * fAttrHiddenSize << ";\n";
602 }
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";
613 }
614 } else {
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";
628 }
629 }
630 out << SP << SP << "}\n";
631
632 // Clip the elements of the update gate and the reset gate into the range [-fClip, fClip]
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";
638 }
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";
643 }
644 out << SP << SP << SP << OpName << "_reset_gate[i] = (r < " << fAttrClip << ") ? r : " << fAttrClip << ";\n";
645 out << SP << SP << "}\n";
646 }
647
648 // Apply the activation function to the update gate and the reset gate
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";
660 }
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";
664 }
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";
686 }
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";
692 }
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";
702 }
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";
708 }
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]
723 << ")\n";
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]
726 << ")\n";
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";
745 } else { // fAttrActivations[direction * 2] = Softplus
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";
750 }
751
752 if (fAttrLinearBeforeReset == 0) {
753 out << SP << SP << "if (seq == 0) {\n";
754 if (!fNInitial_h.empty()) {
755 // feedback = reset_gate o initial_hidden_state
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";
760 }
761 out << SP << SP << "} else {\n";
762 // feedback = reset_gate o previous_hidden_state
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";
767 } else {
768 out << SP << SP << SP << "size_t previous_offset = (seq - 1) * "
769 << num_directions * batch_size * fAttrHiddenSize << ";\n";
770 }
771 } else {
772 out << SP << SP << SP << "size_t previous_offset = (index + 1) * "
773 << num_directions * batch_size * fAttrHiddenSize << " + " << batch_size * fAttrHiddenSize << ";\n";
774 }
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";
780 // feedback = feedback * R_h^T
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";
788 } else { // fAttrLinearBeforeReset=1
789 // feedback = previous_hidden_state * R_h^T
790 // LM fixes
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()) {
796 // feedback = W * initial_hidden_state + bias
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";
801 }
802 out << SP << SP << "} else {\n";
803 // case for seq > 0
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";
808 } else {
809 out << SP << SP << SP << "size_t previous_offset = (seq - 1) * "
810 << num_directions * batch_size * fAttrHiddenSize << ";\n";
811 }
812 } else {
813 out << SP << SP << SP << "size_t previous_offset = (index + 1) * "
814 << num_directions * batch_size * fAttrHiddenSize << " + " << batch_size * fAttrHiddenSize << ";\n";
815 }
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";
820 // endif on seq 0 or not
821 out << SP << SP << "}\n";
822 // Add the bias of the recurrence to feedback
823 if (!fNB.empty()) {
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
828 << "_incy);\n";
829 }
830 // feedback = reset_gate o feedback
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";
834 }
835
836 // hidden_gate = hidden_gate + feedback
837 out << SP << SP << "BLAS::saxpy_(&" << OpName << "_feedback_size, &" << OpName << "_alpha, " << OpName
838 << "_feedback, &" << OpName << "_incx, " << OpName << "_hidden_gate + offset, &" << OpName << "_incy);\n";
839
840 // Clip the elements of the hidden gate into the range [-fClip, fClip]
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";
846 }
847 out << SP << SP << SP << OpName << "_hidden_gate[i] = (x < " << fAttrClip << ") ? x : " << fAttrClip << ";\n";
848 out << SP << SP << "}\n";
849 }
850
851 // Apply the activation function to the hidden gate
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";
861 }
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";
879 }
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";
889 }
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]
901 << ")\n";
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";
915 } else { // fAttrActivations[direction * 2 + 1] = Softplus
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";
919 }
920
921 // hidden_state = (1 - update_gate) o hidden_gate
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";
926
927 out << SP << SP << "if (seq == 0) {\n";
928 if (!fNInitial_h.empty()) {
929 // hidden_state += update_gate o initial_hidden_state
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";
934 }
935 out << SP << SP << "} else {\n";
936 // hidden_state += update_gate o previous_hidden_state
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";
941 } else {
942 out << SP << SP << SP << "size_t previous_offset = (seq - 1) * "
943 << num_directions * batch_size * fAttrHiddenSize << ";\n";
944 }
945 } else {
946 out << SP << SP << SP << "size_t previous_offset = (index + 1) * "
947 << num_directions * batch_size * fAttrHiddenSize << " + " << batch_size * fAttrHiddenSize << ";\n";
948 }
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";
954
955 out << SP << "}\n";
956 }
957
958 // Padding the hidden state for GRU with different sequence lengths
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";
968 out << SP << SP << SP << SP << SP << "}\n";
969 }
970 out << SP << SP << SP << "}\n";
971 out << SP << SP << "}\n";
972 out << SP << "}\n";
973 }
974
975 // Copy the hidden state into y and y_h
976 if (fAttrLayout == 0) {
977 if (!fNY_h.empty()) {
978 // Copy hidden_state into Y_h
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";
984 } else {
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";
988 }
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";
992 }
993 } else { // GRU with different sequence lengths
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";
999 out << SP << "}\n";
1000 } else {
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";
1008 out << SP << "}\n";
1009 }
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
1013 << ";\n";
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";
1018 out << SP << "}\n";
1019 }
1020 }
1021 }
1022 } else { // fAttrLayout=1
1023 if (!fNY.empty()) {
1024 // Copy hidden_state into Y
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";
1035 out << SP << "}\n";
1036 }
1037 }
1038 if (!fNY_h.empty()) {
1039 // Copy the hidden_state into Y_h
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";
1046 out << SP << "}\n";
1047 } else {
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";
1051 } else {
1052 out << SP << SP << "size_t seq = " << "tensor_" << fNSequence_lens << "[batch] - 1;\n";
1053 }
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";
1059 out << SP << "}\n";
1060 }
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
1064 << ";\n";
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";
1069 out << SP << "}\n";
1070 }
1071 }
1072 }
1073
1074 return out.str();
1075}
1076
1077} // namespace TMVA::Experimental::SOFIE
1078
1079#endif
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
Gated Recurrent Unit operator.
std::vector< size_t > fShapeY
Shape of the output.
std::string fNX
Name of the input.
std::string fType
Type of the tensors.
std::string fAttrDirection
Direction of processing.
std::string fNR
Name of the recurrence.
std::vector< float > fAttrActivationBeta
Scaling values used by some activation functions.
std::string fNY
Name of the output.
std::string fNY_h
Name of the last sequence of the output.
std::string fNSequence_lens
Name of the length of the sequences.
std::vector< std::string > fAttrActivations
Activation functions.
void Initialize(RModel &) override
Initialize the model.
ROperator_GRU(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 layout, size_t linear_before_reset, std::string nameX, std::string nameW, std::string nameR, std::string nameB, std::string nameSequence_lens, std::string nameInitial_h, std::string nameY, std::string nameY_h)
Constructor of ROperator_GRU from the attributes.
size_t fAttrHiddenSize
Number of the hidden layers.
std::vector< std::vector< size_t > > ShapeInference(std::vector< std::vector< size_t > >)
Infers the shape of the output tensors.
std::string Generate(std::string) override
Generate the inference code.
std::vector< float > fAttrActivationAlpha
Scaling values used by some activation functions.
std::vector< size_t > fShapeR
Shape of the recurrence.
std::string fNW
Name of the weights.
std::vector< size_t > fShapeX
Shape of the input.
std::vector< size_t > fShapeInitial_h
Shape of the initial value of hidden states.
std::vector< size_t > fShapeSequence_lens
Shape of the length of the sequences.
std::vector< size_t > fShapeY_h
Shape of the last sequence of the output.
size_t fAttrLinearBeforeReset
Linear layer before the reset gate.
std::vector< size_t > fShapeB
Shape of the bias.
std::string fNInitial_h
Name of the initial value of the hidden states.
std::vector< size_t > fShapeW
Shape of the weights.
ROperator_GRU()
Default constructor of ROperator_GRU.
std::vector< std::string > GetBlasRoutines() override
Returns the blas routines needed to compile the generated code.
std::vector< std::string_view > fInputTensorNames
Definition ROperator.hxx:44
std::vector< std::string_view > fOutputTensorNames
Definition ROperator.hxx:45
const Int_t n
Definition legend1.C:16
ETensorType ConvertStringToType(std::string type)