Skip to main content

Module extra

Module extra 

Source
Expand description

Additional activation kernels: Swish, PRelu, LogSoftmax, Hardmax, ThresholdedRelu, Shrink.

Structs§

LogSoftmaxBackward
Backward: dx = dy - softmax(x) * sum(dy)
LogSoftmaxForward
Forward: y = x - log(sum(exp(x))) [numerically stable: subtract max first]
PreluBackward
Backward: dx = dy if x >= 0 else slope * dy; dslope = dy * min(x, 0) = dy * x if x < 0 else 0
PreluForward
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).
ShrinkBackward
Backward: dx = dy if |x| > lambd else 0
ShrinkForward
Forward: y = x - bias if x > lambd, x + bias if x < -lambd, else 0
ShrinkRuntimeOp
RuntimeOp wrapper for Shrink that stores lambd and bias.
SwishBackward
Backward: dx = (sigmoid(x) + x * sigmoid(x) * (1 - sigmoid(x))) * dy = (sig + x * sig * (1 - sig)) * dy
SwishForward
Forward: y = x * sigmoid(x) = x / (1 + exp(-x))
ThresholdedReluBackward
Backward: dx = dy if x > alpha else 0
ThresholdedReluForward
Forward: y = x if x > alpha else 0
ThresholdedReluRuntimeOp
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