File size: 4,914 Bytes
d20d01e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
// pybind11 entry point for the ActionCodec2 physical-stage kernels.
//
// These are the encode/decode kernels the runtime bundled with an ActionCodec2
// Hugging Face artifact calls when a native module is available.  Each one has
// a bit-identical NumPy implementation in that runtime.

#include <pybind11/numpy.h>
#include <pybind11/pybind11.h>

#include <cstdint>
#include <stdexcept>

#include "parallel.hpp"

namespace py = pybind11;

py::array_t<int32_t> so3_second_order_encode_batch(
    const py::array_t<double, py::array::c_style | py::array::forcecast>&,
    const py::array_t<int32_t, py::array::c_style | py::array::forcecast>&,
    const py::array_t<double, py::array::c_style | py::array::forcecast>&,
    const py::array_t<double, py::array::c_style | py::array::forcecast>&,
    int32_t, bool, double, double, bool, int);

py::array_t<int32_t> additive_second_order_encode_batch(
    const py::array_t<float, py::array::c_style | py::array::forcecast>&,
    const py::array_t<int32_t, py::array::c_style | py::array::forcecast>&,
    const py::array_t<double, py::array::c_style | py::array::forcecast>&,
    const py::array_t<double, py::array::c_style | py::array::forcecast>&,
    const py::array_t<float, py::array::c_style | py::array::forcecast>&, bool, int);

py::array_t<float> physical_decode_batch(
    const py::array_t<float>&, const py::array_t<int32_t>&, const py::array_t<uint8_t>&,
    const py::array_t<float>&, const py::array_t<float>&,
    const py::array_t<int32_t, py::array::c_style | py::array::forcecast>&,
    const py::array_t<float, py::array::c_style | py::array::forcecast>&,
    const py::array_t<float, py::array::c_style | py::array::forcecast>&, int32_t, int);

py::array_t<int32_t> physical_encode_batch(
    const py::array_t<float>&, const py::array_t<int32_t>&, const py::array_t<uint8_t>&,
    const py::array_t<float>&, const py::array_t<float>&,
    const py::array_t<float, py::array::c_style | py::array::forcecast>&,
    const py::array_t<float, py::array::c_style | py::array::forcecast>&,
    const py::array_t<float, py::array::c_style | py::array::forcecast>&, int32_t, int,
    const py::array_t<float>&, bool);

py::array_t<int32_t> so3_encode_batch(
    const py::array_t<double, py::array::c_style | py::array::forcecast>&,
    const py::array_t<int32_t, py::array::c_style | py::array::forcecast>&,
    const py::array_t<float, py::array::c_style | py::array::forcecast>&, int32_t, bool, bool,
    double, bool, int);

void register_setbpe(py::module_& module);
void register_resample(py::module_& module);

PYBIND11_MODULE(_actioncodec2_native, module) {
  module.doc() = "ActionCodec2 encode/decode kernels: physical stage, resampling, Set-BPE";
  register_setbpe(module);
  register_resample(module);
  module.def(
      "set_thread_budget",
      [](int threads) {
        if (threads < 0) throw std::invalid_argument("thread budget must be non-negative");
        ac2::thread_budget().store(threads);
      },
      py::arg("threads"),
      "Cap threads per call for this process; 0 uses every CPU it may run on.");
  module.def("thread_budget", [] { return ac2::thread_budget().load(); });
  module.def("plan_threads", &ac2::plan_threads, py::arg("count"), py::arg("requested"),
             py::arg("seconds_each"), "Threads a call would use for ``count`` items.");
  module.def("so3_second_order_encode_batch", &so3_second_order_encode_batch,
             py::arg("table"), py::arg("bins"), py::arg("quaternions"),
             py::arg("previous_first_order"), py::arg("input_order"), py::arg("left"),
             py::arg("limit"), py::arg("radius"), py::arg("smooth"), py::arg("threads") = 0);
  module.def("additive_second_order_encode_batch", &additive_second_order_encode_batch,
             py::arg("table"), py::arg("bins"), py::arg("positions"),
             py::arg("previous_first_order"), py::arg("epsilon"), py::arg("smooth"),
             py::arg("threads") = 0);
  module.def("physical_decode_batch", &physical_decode_batch, py::arg("table"),
             py::arg("bins"), py::arg("binary"), py::arg("lower"), py::arg("upper"),
             py::arg("cells"), py::arg("initial"), py::arg("initial_velocity"),
             py::arg("order"), py::arg("threads") = 0);
  module.def("physical_encode_batch", &physical_encode_batch, py::arg("table"), py::arg("bins"),
             py::arg("binary"), py::arg("lower"), py::arg("upper"), py::arg("positions"),
             py::arg("initial"), py::arg("initial_velocity"), py::arg("order"),
             py::arg("threads") = 0, py::arg("epsilon") = py::array_t<float>(),
             py::arg("smooth") = false);
  module.def("so3_encode_batch", &so3_encode_batch, py::arg("table"), py::arg("bins"),
             py::arg("actions"), py::arg("input_order"), py::arg("rotation_vector"),
             py::arg("left"), py::arg("limit"), py::arg("smooth"), py::arg("threads") = 0);
}