Expand description
Additional activation kernels: Swish, PRelu, LogSoftmax, Hardmax, ThresholdedRelu, Shrink.
Structs§
- LogSoftmax
Backward - Backward: dx = dy - softmax(x) * sum(dy)
- LogSoftmax
Forward - Forward: y = x - log(sum(exp(x))) [numerically stable: subtract max first]
- Prelu
Backward - Backward: dx = dy if x >= 0 else slope * dy; dslope = dy * min(x, 0) = dy * x if x < 0 else 0
- Prelu
Forward - Forward: y = max(0, x) + slope * min(0, x) The slope tensor has the same shape as x (or broadcastable; kernel assumes same shape here).
- Shrink
Backward - Backward: dx = dy if |x| > lambd else 0
- Shrink
Forward - Forward: y = x - bias if x > lambd, x + bias if x < -lambd, else 0
- Shrink
Runtime Op - RuntimeOp wrapper for Shrink that stores lambd and bias.
- Swish
Backward - Backward: dx = (sigmoid(x) + x * sigmoid(x) * (1 - sigmoid(x))) * dy = (sig + x * sig * (1 - sig)) * dy
- Swish
Forward - Forward: y = x * sigmoid(x) = x / (1 + exp(-x))
- Thresholded
Relu Backward - Backward: dx = dy if x > alpha else 0
- Thresholded
Relu Forward - Forward: y = x if x > alpha else 0
- Thresholded
Relu Runtime Op - RuntimeOp wrapper for ThresholdedRelu that stores the alpha scalar.
Functions§
- log_
softmax_ backward - Backward: dx = dy - softmax(x) * sum(dy)
- log_
softmax_ forward - Forward: y = x - log(sum(exp(x))) [numerically stable: subtract max first]
- prelu_
backward - Backward: dx = dy if x >= 0 else slope * dy; dslope = dy * min(x, 0) = dy * x if x < 0 else 0
- prelu_
forward - Forward: y = max(0, x) + slope * min(0, x) The slope tensor has the same shape as x (or broadcastable; kernel assumes same shape here).
- shrink_
backward - Backward: dx = dy if |x| > lambd else 0
- shrink_
forward - Forward: y = x - bias if x > lambd, x + bias if x < -lambd, else 0
- swish_
backward - Backward: dx = (sigmoid(x) + x * sigmoid(x) * (1 - sigmoid(x))) * dy = (sig + x * sig * (1 - sig)) * dy
- swish_
forward - Forward: y = x * sigmoid(x) = x / (1 + exp(-x))
- thresholded_
relu_ backward - Backward: dx = dy if x > alpha else 0
- thresholded_
relu_ forward - Forward: y = x if x > alpha else 0