MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / test_lstm

Function test_lstm

dnn/test/arm_common/lstm.cpp:79–122  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

77
78namespace {
79void test_lstm(bool bias, bool direction, Handle* handle) {
80 Checker<LSTM> checker(handle, true);
81 //! because lstm has tanh, exp mathematical compute, after more iteration,
82 //! the error will more than 1e-3
83 checker.set_epsilon(1e-2);
84 checker.set_output_canonizer(output_canonizer);
85 for (size_t input_size : {2, 8, 13})
86 for (size_t hidden_size : {1, 4, 17}) {
87 size_t dir_size = direction == false ? 1 : 2;
88 LSTM::Param param;
89 param.bidirectional = direction;
90 size_t gate_hidden_size = 4 * hidden_size;
91 param.bias = bias;
92 param.hidden_size = hidden_size;
93 for (size_t seq_len : {1, 3, 5})
94 for (size_t batch_size : {1, 2, 4})
95 for (size_t number_layer : {1, 2, 4, 5, 8}) {
96 size_t flatten_size = 0;
97 for (size_t layer = 0; layer < number_layer; layer++) {
98 for (size_t dir = 0; dir < dir_size; dir++) {
99 flatten_size += layer == 0
100 ? input_size
101 : dir_size * hidden_size; // ih
102 flatten_size += hidden_size; // hh
103 }
104 }
105 if (bias) {
106 flatten_size += 2 * dir_size * number_layer;
107 }
108 param.num_layers = number_layer;
109 checker.set_param(param).exec(
110 {{seq_len, batch_size, input_size}, // input
111 {number_layer * dir_size, batch_size,
112 hidden_size}, // hx
113 {number_layer * dir_size, batch_size,
114 hidden_size}, // hy
115 {gate_hidden_size, flatten_size}, // flat weight
116 {},
117 {},
118 {},
119 {}});
120 }
121 }
122}
123
124} // namespace
125

Callers 1

TEST_FFunction · 0.70

Calls 1

execMethod · 0.45

Tested by

no test coverage detected