TensorPlay AI

自定义自动微分

当标准算子组合不够高效、稳定,或需要连接自定义硬件时,直接定义梯度逻辑。

定义 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 标记无需梯度的输出。
Ask DeepWiki