| 2135 | } |
| 2136 | |
| 2137 | void elu_gradient ( |
| 2138 | tensor& grad, |
| 2139 | const tensor& dest, |
| 2140 | const tensor& gradient_input, |
| 2141 | const float alpha |
| 2142 | ) |
| 2143 | { |
| 2144 | const auto out = grad.host(); |
| 2145 | const auto in = dest.host(); |
| 2146 | const auto gi = gradient_input.host(); |
| 2147 | if (is_same_object(grad, gradient_input)) |
| 2148 | { |
| 2149 | for (size_t i = 0; i < dest.size(); ++i) |
| 2150 | { |
| 2151 | if (in[i] > 0) |
| 2152 | out[i] = gi[i]; |
| 2153 | else |
| 2154 | out[i] = (alpha + in[i]) * gi[i]; |
| 2155 | } |
| 2156 | } |
| 2157 | else |
| 2158 | { |
| 2159 | for (size_t i = 0; i < dest.size(); ++i) |
| 2160 | { |
| 2161 | if (in[i] > 0) |
| 2162 | out[i] += gi[i]; |
| 2163 | else |
| 2164 | out[i] += (alpha + in[i]) * gi[i]; |
| 2165 | } |
| 2166 | } |
| 2167 | } |
| 2168 | |
| 2169 | // ---------------------------------------------------------------------------------------- |
| 2170 | |