Skip to main content

Module elemwise_binary

Module elemwise_binary 

Source

Structs§

ClipRuntimeOp
A RuntimeOp wrapper for Clip that stores the min/max scalar params alongside the kernel struct (which only stores block_size).
ElemwiseClipBackward
Backward: pass dy through only where x was in [min_val, max_val]
ElemwiseClipForward
Forward: out = clamp(x, min_val, max_val)
ElemwiseDivBackward
Backward: da = dy / b, db = -a * dy / b^2
ElemwiseDivForward
Forward: out = a / b
ElemwiseEqualForward
Forward: out = 1.0 if a == b else 0.0
ElemwiseFmodForward
Forward fmod: out = a - trunc(a/b)*b (C-style float remainder)
ElemwiseGreaterEqualForward
Forward: out = 1.0 if a >= b else 0.0
ElemwiseGreaterForward
Forward: out = 1.0 if a > b else 0.0
ElemwiseLessEqualForward
Forward: out = 1.0 if a <= b else 0.0
ElemwiseLessForward
Forward: out = 1.0 if a < b else 0.0
ElemwiseMaxBackward
Backward: pass dy to the input that was larger, 0 to the other.
ElemwiseMaxForward
Forward: out = max(a, b)
ElemwiseMeanBackward
Backward: da = db = dy / 2
ElemwiseMeanForward
Forward: out = (a + b) / 2
ElemwiseMinBackward
Backward: pass dy to the input that was smaller, 0 to the other.
ElemwiseMinForward
Forward: out = min(a, b)
ElemwiseMulBackward
Backward: da = dy * b, db = dy * a
ElemwiseMulForward
Forward: out = a * b
ElemwisePowBackward
Backward: da = b * a^(b-1) * dy, db = log(a) * a^b * dy
ElemwisePowForward
Forward: out = a ^ b = exp(b * log(a))
ElemwiseSubBackward
Backward: da = dy, db = -dy
ElemwiseSubForward
Forward: out = a - b
ElemwiseSumBackward
Backward: da = db = dy
ElemwiseSumForward
Forward: out = a + b (binary ElemSum)
ElemwiseWhereBackward
Backward: dx = where(cond, dy, 0), dy_in = where(cond, 0, dy)
ElemwiseWhereForward
Forward: out = x where cond != 0 else y

Functions§

elemwise_clip_backward
Backward: pass dy through only where x was in [min_val, max_val]
elemwise_clip_forward
Forward: out = clamp(x, min_val, max_val)
elemwise_div_backward
Backward: da = dy / b, db = -a * dy / b^2
elemwise_div_forward
Forward: out = a / b
elemwise_equal_forward
Forward: out = 1.0 if a == b else 0.0
elemwise_fmod_forward
Forward fmod: out = a - trunc(a/b)*b (C-style float remainder)
elemwise_greater_equal_forward
Forward: out = 1.0 if a >= b else 0.0
elemwise_greater_forward
Forward: out = 1.0 if a > b else 0.0
elemwise_less_equal_forward
Forward: out = 1.0 if a <= b else 0.0
elemwise_less_forward
Forward: out = 1.0 if a < b else 0.0
elemwise_max_backward
Backward: pass dy to the input that was larger, 0 to the other.
elemwise_max_forward
Forward: out = max(a, b)
elemwise_mean_backward
Backward: da = db = dy / 2
elemwise_mean_forward
Forward: out = (a + b) / 2
elemwise_min_backward
Backward: pass dy to the input that was smaller, 0 to the other.
elemwise_min_forward
Forward: out = min(a, b)
elemwise_mul_backward
Backward: da = dy * b, db = dy * a
elemwise_mul_forward
Forward: out = a * b
elemwise_pow_backward
Backward: da = b * a^(b-1) * dy, db = log(a) * a^b * dy
elemwise_pow_forward
Forward: out = a ^ b = exp(b * log(a))
elemwise_sub_backward
Backward: da = dy, db = -dy
elemwise_sub_forward
Forward: out = a - b
elemwise_sum_backward
Backward: da = db = dy
elemwise_sum_forward
Forward: out = a + b (binary ElemSum)
elemwise_where_backward
Backward: dx = where(cond, dy, 0), dy_in = where(cond, 0, dy)
elemwise_where_forward
Forward: out = x where cond != 0 else y