| 12 | namespace py = pybind11; |
| 13 | |
| 14 | kernels::Kernel* parse_kernel_spec (const py::object& kernel_spec) { |
| 15 | |
| 16 | if (!py::hasattr(kernel_spec, "is_kernel")) throw std::invalid_argument("invalid kernel"); |
| 17 | |
| 18 | // Deal with operators first |
| 19 | bool is_kernel = py::bool_(kernel_spec.attr("is_kernel")); |
| 20 | if (!is_kernel) { |
| 21 | kernels::Kernel *k1, *k2; |
| 22 | py::object spec1 = kernel_spec.attr("k1"), |
| 23 | spec2 = kernel_spec.attr("k2"); |
| 24 | k1 = parse_kernel_spec(spec1); |
| 25 | k2 = parse_kernel_spec(spec2); |
| 26 | if (k1->get_ndim() != k2->get_ndim()) throw dimension_mismatch(); |
| 27 | size_t op = py::int_(kernel_spec.attr("operator_type")); |
| 28 | if (op == 0) { |
| 29 | return new kernels::Sum(k1, k2); |
| 30 | } else if (op == 1) { |
| 31 | return new kernels::Product(k1, k2); |
| 32 | } else { |
| 33 | throw std::invalid_argument("unrecognized operator"); |
| 34 | } |
| 35 | } |
| 36 | |
| 37 | |
| 38 | kernels::Kernel* kernel; |
| 39 | size_t kernel_type = py::int_(kernel_spec.attr("kernel_type")); |
| 40 | switch (kernel_type) { |
| 41 | {% for spec in specs %} |
| 42 | case {{ spec.index }}: { |
| 43 | {% if spec.stationary %} |
| 44 | py::object metric = kernel_spec.attr("metric"); |
| 45 | size_t metric_type = py::int_(metric.attr("metric_type")); |
| 46 | size_t ndim = py::int_(metric.attr("ndim")); |
| 47 | py::list axes = py::list(metric.attr("axes")); |
| 48 | bool blocked = py::bool_(kernel_spec.attr("blocked")); |
| 49 | py::array_t<double> min_block = py::array_t<double>(kernel_spec.attr("min_block")); |
| 50 | py::array_t<double> max_block = py::array_t<double>(kernel_spec.attr("max_block")); |
| 51 | |
| 52 | // Select the correct template based on the metric type |
| 53 | if (metric_type == 0) { |
| 54 | kernel = new kernels::{{ spec.name }}<metrics::IsotropicMetric> ( |
| 55 | {% for param in spec.params %} |
| 56 | py::float_(kernel_spec.attr("{{ param }}")), |
| 57 | {%- endfor %} |
| 58 | {% for con in spec.constants %} |
| 59 | py::float_(kernel_spec.attr("{{ con.name }}")), |
| 60 | {%- endfor %} |
| 61 | blocked, |
| 62 | (double*)&(min_block.unchecked<1>()(0)), |
| 63 | (double*)&(max_block.unchecked<1>()(0)), |
| 64 | ndim, |
| 65 | py::len(axes) |
| 66 | ); |
| 67 | } else if (metric_type == 1) { |
| 68 | kernel = new kernels::{{ spec.name }}<metrics::AxisAlignedMetric> ( |
| 69 | {% for param in spec.params %} |
| 70 | py::float_(kernel_spec.attr("{{ param }}")), |
| 71 | {%- endfor %} |
nothing calls this directly
no test coverage detected