MCPcopy Create free account
hub / github.com/ml-explore/mlx / CustomTransforms

Method CustomTransforms

mlx/primitives.h:836–854  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

834class CustomTransforms : public Primitive {
835 public:
836 explicit CustomTransforms(
837 Stream stream,
838 int num_outputs,
839 std::function<std::vector<array>(
840 const std::vector<array>&,
841 const std::vector<array>&,
842 const std::vector<array>&)> vjp,
843 std::function<std::vector<array>(
844 const std::vector<array>&,
845 const std::vector<array>&,
846 const std::vector<int>&)> jvp,
847 std::function<std::pair<std::vector<array>, std::vector<int>>(
848 const std::vector<array>&,
849 const std::vector<int>&)> vmap)
850 : Primitive(stream),
851 num_outputs_(num_outputs),
852 vjp_fun_(std::move(vjp)),
853 jvp_fun_(std::move(jvp)),
854 vmap_fun_(std::move(vmap)) {}
855
856 void eval_cpu(const std::vector<array>& inputs, std::vector<array>& outputs)
857 override;

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected