Skip to main contentModule elemwise_binary
Source - 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
- 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