自定义自动微分
当标准算子组合不够高效、稳定,或需要连接自定义硬件时,直接定义梯度逻辑。
定义 Function
import tensorplay as tp
from tensorplay.autograd import Function
class MyExp(Function):
@staticmethod
def forward(ctx, value):
result = value.exp()
ctx.save_for_backward(result)
return result
@staticmethod
def backward(ctx, grad_output):
result, = ctx.saved_tensors
return grad_output * result
x = tp.randn(3, requires_grad=True)
y = MyExp.apply(x)
y.sum().backward()何时使用
- 合并梯度计算以提升效率。
- 实现更稳定的数值形式。
- 连接带有专用微分逻辑的硬件。
- 为量化等不可微操作提供替代梯度。
ctx 上下文
- save_for_backward 保存反向传播所需张量。
- saved_tensors 读取已保存张量。
- mark_dirty 标记原地修改。
- mark_non_differentiable 标记无需梯度的输出。
