/usr/local/lib64/python3.6/site-packages/torch/include/torch/csrc/autograd/utils
Edit: /usr/local/lib64/python3.6/site-packages/torch/include/torch/csrc/autograd/utils/lambda_post_hook.h (976B)
#pragma once
#include
namespace torch {
namespace autograd {
namespace utils {
// Turns lambda into a torch::autograd::FunctionPostHook.
class LambdaPostHook : public torch::autograd::FunctionPostHook {
using variable_list = std::vector;
public:
// The lambda function takes as arguments the outputs and inputs of the
// autograd function and can modify the outputs of the autograd function by
// returning a new output if needed.
/* implicit */ LambdaPostHook(
std::function
fn)
: fn_(std::move(fn)) {}
variable_list operator()(
const variable_list& outputs,
const variable_list& inputs) override {
return fn_(outputs, inputs);
}
protected:
std::function fn_;
};
} // namespace utils
} // namespace autograd
} // namespace torch