// Intrinsic first- and stateful second-order SO(3) error-feedback encoders. #include #include #include #include #include #include #include #include #include "parallel.hpp" namespace py = pybind11; namespace { struct Quaternion { double x; double y; double z; double w; }; Quaternion normalize(Quaternion q) { const double norm = std::sqrt(q.x * q.x + q.y * q.y + q.z * q.z + q.w * q.w); q.x /= norm; q.y /= norm; q.z /= norm; q.w /= norm; return q; } Quaternion multiply(const Quaternion& a, const Quaternion& b) { return normalize({a.w * b.x + a.x * b.w + a.y * b.z - a.z * b.y, a.w * b.y - a.x * b.z + a.y * b.w + a.z * b.x, a.w * b.z + a.x * b.y - a.y * b.x + a.z * b.w, a.w * b.w - a.x * b.x - a.y * b.y - a.z * b.z}); } Quaternion conjugate(const Quaternion& q) { return {-q.x, -q.y, -q.z, q.w}; } Quaternion from_euler(const float* euler) { const double roll = 0.5 * static_cast(euler[0]); const double pitch = 0.5 * static_cast(euler[1]); const double yaw = 0.5 * static_cast(euler[2]); const double sr = std::sin(roll), cr = std::cos(roll); const double sp = std::sin(pitch), cp = std::cos(pitch); const double sy = std::sin(yaw), cy = std::cos(yaw); return normalize({sr * cp * cy - cr * sp * sy, cr * sp * cy + sr * cp * sy, cr * cp * sy - sr * sp * cy, cr * cp * cy + sr * sp * sy}); } void to_rotvec(Quaternion q, double* output) { if (q.w < 0.0) { q.x = -q.x; q.y = -q.y; q.z = -q.z; q.w = -q.w; } const double norm = std::sqrt(q.x * q.x + q.y * q.y + q.z * q.z); const double scale = norm < 1e-12 ? 2.0 : 2.0 * std::atan2(norm, std::max(0.0, q.w)) / norm; output[0] = q.x * scale; output[1] = q.y * scale; output[2] = q.z * scale; } Quaternion from_rotvec(const double* value) { const double angle = std::sqrt(value[0] * value[0] + value[1] * value[1] + value[2] * value[2]); const double scale = angle < 1e-12 ? 0.5 : std::sin(0.5 * angle) / angle; return normalize({value[0] * scale, value[1] * scale, value[2] * scale, std::cos(0.5 * angle)}); } Quaternion from_input(const float* value, bool rotation_vector) { if (!rotation_vector) return from_euler(value); const double converted[3] = {static_cast(value[0]), static_cast(value[1]), static_cast(value[2])}; return from_rotvec(converted); } Quaternion residual(const Quaternion& target, const Quaternion& reconstruction, bool left) { return left ? multiply(target, conjugate(reconstruction)) : multiply(conjugate(reconstruction), target); } Quaternion apply(const Quaternion& state, const Quaternion& increment, bool left) { return left ? multiply(increment, state) : multiply(state, increment); } std::array as_rotvec(const Quaternion& value) { std::array output{}; to_rotvec(value, output.data()); return output; } struct SO3Layout { const double* table; const int32_t* bins; int64_t width; double limit; bool smooth; bool rotation_vector; bool left; int32_t input_order; double projection_radius = std::numeric_limits::infinity(); }; void project_primitive(double* value, double radius) { if (!std::isfinite(radius)) return; const double norm = std::sqrt(value[0]*value[0] + value[1]*value[1] + value[2]*value[2]); if (norm > radius) { const double scale = radius / norm; for (int j = 0; j < 3; ++j) value[j] *= scale; } } inline double primitive(const SO3Layout& layout, int axis, int32_t index) { return layout.table[static_cast(axis) * layout.width + index]; } std::array nearest_choice(const SO3Layout& layout, const std::array& request) { std::array choice{}; for (int axis = 0; axis < 3; ++axis) { double best = std::numeric_limits::infinity(); for (int32_t index = 0; index < layout.bins[axis]; ++index) { const double distance = std::fabs(request[axis] - primitive(layout, axis, index)); if (distance < best) { best = distance; choice[axis] = index; } } } return choice; } std::vector smooth_axis(const SO3Layout& layout, int axis, double request, int32_t nearest) { std::vector output; bool has_nearest = false; const double radius = 2.0 * layout.limit + 1e-12; for (int32_t index = 0; index < layout.bins[axis]; ++index) { if (std::fabs(primitive(layout, axis, index) - request) <= radius) { output.push_back(index); has_nearest = has_nearest || index == nearest; } } if (!has_nearest) output.push_back(nearest); return output; } bool earlier(double error, const std::array& choice, double best_error, const std::array& best_choice, bool found) { return !found || error < best_error || (error == best_error && choice < best_choice); } std::array smooth_choice(const SO3Layout& layout, const std::array& request, const Quaternion& reconstruction, const Quaternion& target, const std::array& direction, const std::array& nearest) { const auto x = smooth_axis(layout, 0, request[0], nearest[0]); const auto y = smooth_axis(layout, 1, request[1], nearest[1]); const auto z = smooth_axis(layout, 2, request[2], nearest[2]); std::array best_any{}, best_preferred{}; double any_error = 0.0, preferred_error = 0.0; bool found_any = false, found_preferred = false; for (const int32_t ix : x) { for (const int32_t iy : y) { for (const int32_t iz : z) { const std::array choice{ix, iy, iz}; double vector[3] = {primitive(layout, 0, ix), primitive(layout, 1, iy), primitive(layout, 2, iz)}; project_primitive(vector, layout.projection_radius); const Quaternion rebuilt = apply(reconstruction, from_rotvec(vector), layout.left); const auto remaining = as_rotvec(residual(target, rebuilt, layout.left)); const double error = std::sqrt(remaining[0] * remaining[0] + remaining[1] * remaining[1] + remaining[2] * remaining[2]); if (error > layout.limit + 1e-12) continue; if (earlier(error, choice, any_error, best_any, found_any)) { best_any = choice; any_error = error; found_any = true; } const double dot = vector[0] * direction[0] + vector[1] * direction[1] + vector[2] * direction[2]; if (dot >= -1e-15 && earlier(error, choice, preferred_error, best_preferred, found_preferred)) { best_preferred = choice; preferred_error = error; found_preferred = true; } } } } return found_preferred ? best_preferred : (found_any ? best_any : nearest); } void encode_trajectory(const SO3Layout& layout, const float* actions, int64_t steps, int32_t* cells) { Quaternion truth{0.0, 0.0, 0.0, 1.0}; Quaternion reconstruction = truth; std::array direction{}; for (int64_t step = 0; step < steps; ++step) { const Quaternion value = from_input(actions + 3 * step, layout.rotation_vector); truth = layout.input_order == 0 ? value : apply(truth, value, layout.left); const auto request = as_rotvec(residual(truth, reconstruction, layout.left)); auto choice = nearest_choice(layout, request); if (layout.smooth) { choice = smooth_choice(layout, request, reconstruction, truth, direction, choice); } const double vector[3] = {primitive(layout, 0, choice[0]), primitive(layout, 1, choice[1]), primitive(layout, 2, choice[2])}; reconstruction = apply(reconstruction, from_rotvec(vector), layout.left); if (vector[0] != 0.0 || vector[1] != 0.0 || vector[2] != 0.0) { direction = {vector[0], vector[1], vector[2]}; } for (int axis = 0; axis < 3; ++axis) cells[3 * step + axis] = choice[axis]; } } } // namespace py::array_t so3_second_order_encode_batch( const py::array_t& table, const py::array_t& bins, const py::array_t& quaternions, const py::array_t& previous_first_order, int32_t input_order, bool left, double limit, double radius, bool smooth, int threads) { if (table.ndim() != 2 || table.shape(0) != 3 || bins.size() != 3 || quaternions.ndim() != 3 || quaternions.shape(2) != 4 || previous_first_order.ndim() != 2 || previous_first_order.shape(1) != 4 || previous_first_order.shape(0) != quaternions.shape(0) || (input_order != 0 && input_order != 1) || !(limit > 0) || !(radius > 0)) { throw std::invalid_argument("invalid second-order SO(3) batch arguments"); } for (int j = 0; j < 3; ++j) { if (bins.data()[j] < 1 || bins.data()[j] > table.shape(1)) throw std::invalid_argument("SO(3) primitive count exceeds table width"); } const SO3Layout layout{table.data(), bins.data(), table.shape(1), limit, smooth, true, left, input_order, radius}; const int64_t batch = quaternions.shape(0), steps = quaternions.shape(1); py::array_t output({batch, steps, int64_t{3}}); const double* values = quaternions.data(); const double* previous = previous_first_order.data(); int32_t* result = output.mutable_data(); const Quaternion identity{0, 0, 0, 1}; py::gil_scoped_release release; ac2::parallel_for(batch, threads, [&](int64_t i) { Quaternion truth = identity, reconstruction = identity; Quaternion increment{previous[4*i], previous[4*i+1], previous[4*i+2], previous[4*i+3]}; std::array direction{}; for (int64_t t = 0; t < steps; ++t) { const double* raw = values + (i*steps+t)*4; const Quaternion value{raw[0], raw[1], raw[2], raw[3]}; truth = input_order == 0 ? value : apply(truth, value, left); const auto desired = residual(truth, reconstruction, left); const auto correction = residual(desired, increment, left); const auto request = as_rotvec(correction); auto choice = nearest_choice(layout, request); if (smooth) { // Python scores U^-1 W for both conventions; use exactly that order. SO3Layout selection = layout; selection.left = false; choice = smooth_choice(selection, request, identity, correction, direction, choice); } double vector[3] = {primitive(layout, 0, choice[0]), primitive(layout, 1, choice[1]), primitive(layout, 2, choice[2])}; project_primitive(vector, radius); increment = apply(increment, from_rotvec(vector), left); reconstruction = apply(reconstruction, increment, left); if (vector[0] != 0 || vector[1] != 0 || vector[2] != 0) direction = {vector[0], vector[1], vector[2]}; for (int j = 0; j < 3; ++j) result[(i*steps+t)*3+j] = choice[j]; } }); return output; } py::array_t so3_encode_batch( const py::array_t& table, const py::array_t& bins, const py::array_t& actions, int32_t input_order, bool rotation_vector, bool left, double limit, bool smooth, int threads) { if (table.ndim() != 2 || table.shape(0) != 3 || bins.size() != 3 || actions.ndim() != 3 || actions.shape(2) != 3 || (input_order != 0 && input_order != 1) || limit <= 0.0) { throw std::invalid_argument("invalid SO(3) batch encoder arguments"); } for (int axis = 0; axis < 3; ++axis) { if (bins.data()[axis] < 1 || bins.data()[axis] > table.shape(1)) { throw std::invalid_argument("SO(3) primitive count exceeds table width"); } } const SO3Layout layout{table.data(), bins.data(), table.shape(1), limit, smooth, rotation_vector, left, input_order}; const int64_t batch = actions.shape(0); const int64_t steps = actions.shape(1); py::array_t output({batch, steps, int64_t{3}}); { py::gil_scoped_release release; ac2::parallel_for(batch, threads, [&](int64_t index) { encode_trajectory(layout, actions.data() + index * steps * 3, steps, output.mutable_data() + index * steps * 3); }); } return output; }