eagle.frameworks.torch.function#
- eagle.frameworks.torch.function(kernel, *, wrt=None, cache_dir=None, device_arch=None)[source]#
Wrap
kernelas a torch-differentiable callable.- Parameters:
kernel – a per-sample hawk kernel (
@hawk.kernel). Its inputs may beScalar/Vector/Matrixplanes,Paramuniforms and aTerminatedmask; its outputs areMutableplanes.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: