eagle.frameworks.torch.function

Contents

eagle.frameworks.torch.function#

eagle.frameworks.torch.function(kernel, *, wrt=None, cache_dir=None, device_arch=None)[source]#

Wrap kernel as a torch-differentiable callable.

Parameters:
  • kernel – a per-sample hawk kernel (@hawk.kernel). Its inputs may be Scalar/Vector/Matrix planes, Param uniforms and a Terminated mask; its outputs are Mutable planes.

  • wrt – names to differentiate (default: every differentiable input).

  • cache_dir – where artifacts publish (default: a private temp dir).

  • device_arch – CUDA arch to compile for (default: the first call’s device).

Returns:

a KernelFunction.