Repository navigation
Add vectorised select(), any(), and all() functions - #2853
Conversation
|
I've also added the |
|
Ping me when this is ready for review! |
Thanks @SteveBronder! This should be good to look at now |
|
@SteveBronder one more thing I've been thinking about, do we want to start requiring/recommending tuple compatibility for utility functions like these? (probably only really relevant for |
|
I'm not sure if it should be a requirement, but it would def be nice! |
|
@SteveBronder just addressed your comments. I've also changed the handling for tuples in |
| template <typename T_true, typename T_false, | ||
| require_all_stan_scalar_t<T_true, T_false>* = nullptr> | ||
| inline auto select(const bool c, const T_true y_true, const T_false y_false) { | ||
| return c ? y_true : y_false; |
There was a problem hiding this comment.
Should we cast y_true / y_false to the same type explicitly or is it okay if we assume they are implicitly promotable?
There was a problem hiding this comment.
I think we want to do this implicit vs. explicit promotion everywhere
| if (c) { | ||
| return T_true_plain(y_true); | ||
| } else { | ||
| return T_true_plain(y_false); | ||
| } |
There was a problem hiding this comment.
| if (c) { | |
| return T_true_plain(y_true); | |
| } else { | |
| return T_true_plain(y_false); | |
| } | |
| return c ? T_true_plain(y_true) : T_true_plain(y_false); |
| template <typename T_bool, typename T_true, typename T_false, | ||
| require_eigen_array_t<T_bool>* = nullptr, | ||
| require_all_stan_scalar_t<T_true, T_false>* = nullptr> |
There was a problem hiding this comment.
Should we be more restrictive with T_bool here like require_eigen_array_vt<is_integral, T_bool> (bool is integrable here)
SteveBronder
left a comment
There was a problem hiding this comment.
sry clicked approved by accident. Few qs above!
|
I'll take a look at this today |
SteveBronder
left a comment
There was a problem hiding this comment.
Mostly doc fixes but once that's done then I think good to merge!
| * @param x boolean input | ||
| * @return The input unchanged | ||
| */ | ||
| inline bool all(bool x) { return x; } |
There was a problem hiding this comment.
(optional) constexpr?
There was a problem hiding this comment.
Also should this be templated to accept any integral type? Not sure if we would get warnings about downcasting? This applies to all the bool as input functions
| * | ||
| * Overload for Eigen types | ||
| * | ||
| * @tparam Eigen type of the input |
There was a problem hiding this comment.
| * @tparam Eigen type of the input | |
| * @tparam ContainerT A type derived from `Eigen::EigenBase` that has an `integral` scalar type |
Double check the docs for all these
| * Return the second argument if the first argument is true | ||
| * and otherwise return the third argument. | ||
| * | ||
| * <code>select(c, y1, y0) = c ? y1 : y0</code>. |
There was a problem hiding this comment.
(optional) I just prefer using ticks here ` doxygen is able to parse that like it's markdown
| template < | ||
| typename T_true, typename T_false, | ||
| typename T_return = return_type_t<T_true, T_false>, | ||
| typename T_true_plain = promote_scalar_t<T_return, plain_type_t<T_true>>, | ||
| typename T_false_plain = promote_scalar_t<T_return, plain_type_t<T_false>>, | ||
| require_all_container_t<T_true, T_false>* = nullptr, | ||
| require_all_same_t<T_true_plain, T_false_plain>* = nullptr> | ||
| inline T_true_plain select(const bool c, const T_true y_true, | ||
| const T_false y_false) { | ||
| check_matching_dims("select", "y_true", y_true, "y_false", y_false); | ||
| return c ? T_true_plain(y_true) : T_true_plain(y_false); | ||
| } |
There was a problem hiding this comment.
(optional) it may be nice to have a signature for when T_true_plain is not constructable from a T_false_plain that has a static assert in it with a nice message.
| inline T_true_plain select(const bool c, const T_true y_true, | ||
| const T_false y_false) { | ||
| check_matching_dims("select", "y_true", y_true, "y_false", y_false); | ||
| return c ? T_true_plain(y_true) : T_true_plain(y_false); | ||
| } |
There was a problem hiding this comment.
If these are containers should they be passed as references like
| inline T_true_plain select(const bool c, const T_true y_true, | |
| const T_false y_false) { | |
| check_matching_dims("select", "y_true", y_true, "y_false", y_false); | |
| return c ? T_true_plain(y_true) : T_true_plain(y_false); | |
| } | |
| inline T_true_plain select(const bool c, T_true&& y_true, | |
| T_false&& y_false) { | |
| check_matching_dims("select", "y_true", y_true, "y_false", y_false); | |
| return c ? T_true_plain(std::forward<T_true>(y_true)) : T_true_plain(std::forward<T_false>(y_false)); | |
| } |
| require_eigen_array_vt<std::is_integral, T_bool>* = nullptr, | ||
| require_all_stan_scalar_t<T_true, T_false>* = nullptr> | ||
| inline auto select(const T_bool c, const T_true y_true, const T_false y_false) { | ||
| return c.unaryExpr([&](bool cond) { return cond ? y_true : y_false; }).eval(); |
There was a problem hiding this comment.
(optional) I personally prefer we just pull in the names of the things we want into a lambda instead of using &
| return c.unaryExpr([&](bool cond) { return cond ? y_true : y_false; }).eval(); | |
| return c.unaryExpr([y_true, y_false](bool cond) { return cond ? y_true : y_false; }).eval(); |
| check_consistent_sizes("select", "boolean", c, "y_true", y_true, "y_false", | ||
| y_false); |
There was a problem hiding this comment.
(optional) Should the names here be left hand side and right hand side? Like can this error be propagated up to users of Stan?
|
@SteveBronder this is ready for another look |
SteveBronder
left a comment
There was a problem hiding this comment.
Mostly Qs on the signatures for containers and scalars as well as some docs.
| template <typename T, require_integral_t<T>* = nullptr> | ||
| constexpr inline T all(T x) { | ||
| return x; | ||
| } |
There was a problem hiding this comment.
We still want a bool return type right?
| template <typename T, require_integral_t<T>* = nullptr> | |
| constexpr inline T all(T x) { | |
| return x; | |
| } | |
| template <typename T, require_integral_t<T>* = nullptr> | |
| constexpr inline bool all(T x) { | |
| return x; | |
| } |
| * | ||
| * `select(c, y1, y0) = c ? y1 : y0`. | ||
| * | ||
| * @tparam T_true type of the true argument |
There was a problem hiding this comment.
| * @tparam T_true type of the true argument | |
| * @tparam T_true A stan `Scalar` type |
Generally for overloads try to have the type docs reflect the specific types for this overload
There was a problem hiding this comment.
Please fix this for all the docs for the types
| * Return the second argument if the first argument is true | ||
| * and otherwise return the third argument. Eigen expressions are |
There was a problem hiding this comment.
I'd rephrase these slightly so it goes from left to right in terms of arguments
| * Return the second argument if the first argument is true | |
| * and otherwise return the third argument. Eigen expressions are | |
| * If first argument is true return the second argument, else return the third argument |
| return c | ||
| .unaryExpr( | ||
| [y_true, y_false](bool cond) { return cond ? y_true : y_false; }) | ||
| .eval(); |
There was a problem hiding this comment.
How is the return type decided here? I think we should do like the other functions and explicitly cast to our intended return type
| return c | |
| .unaryExpr( | |
| [y_true, y_false](bool cond) { return cond ? y_true : y_false; }) | |
| .eval(); | |
| using ret_t = return_type_t<T_true, T_false>; | |
| return c | |
| .unaryExpr( | |
| [y_true, y_false](bool cond) { return cond ? ret_t(y_true) : ret_t(y_false); }) | |
| .eval(); |
| inline auto select(const T_bool c, const T_true y_true, const T_false y_false) { | ||
| check_consistent_sizes("select", "boolean", c, "left hand side", y_true, | ||
| "right hand side", y_false); | ||
| return c.select(y_true, y_false).eval(); |
There was a problem hiding this comment.
Same as before, I think we should be casting y_true and y_false to the same type before doing the select. idk how Eigen handles that under the hood and I'd rather just be explicit about it
|
@SteveBronder ready for another look |
SteveBronder
left a comment
There was a problem hiding this comment.
Some very minor comments on the docs then I think this is good to go!
| * @tparam Type of container | ||
| * @param x Nested container of boolean inputs |
There was a problem hiding this comment.
docs do not match signature
| * Overload for a single boolean input | ||
| * | ||
| * @tparam T The type of integral input. | ||
| * @param x boolean input |
There was a problem hiding this comment.
| * @param x boolean input | |
| * @param x integral input |
| * | ||
| * Overload for a single boolean input | ||
| * | ||
| * @tparam T The type of integral input. |
There was a problem hiding this comment.
| * @tparam T The type of integral input. | |
| * @tparam T Any type convertible to `bool` |
Summary
This PR extracts the vectorised
select()function introduced in this PR (which in turn is adding the currently OpenCL-only function to the rest of the Math library), and also introduces theany()andall()helper functions.The
select()function is used for ternary operations that are intended to be agnostic between scalars and containers. If the function is used with a mixture of scalars and containers, then the scalar is promoted to the same size and type as the container.Tests
Tests are added for all combinations of scalar/container inputs, as well as error checking for mismatched container sizes.
Side Effects
N/A
Release notes
Added
select()function for vectorised ternary operations, as well as theany()andall()boolean reduction functionsChecklist
Math issue Add
select()andany()helper functions #2852Copyright holder: Andrew Johnson
The copyright holder is typically you or your assignee, such as a university or company. By submitting this pull request, the copyright holder is agreeing to the license the submitted work under the following licenses:
- Code: BSD 3-clause (https://opensource.org/licenses/BSD-3-Clause)
- Documentation: CC-BY 4.0 (https://creativecommons.org/licenses/by/4.0/)
the basic tests are passing
./runTests.py test/unit)make test-headers)make test-math-dependencies)make doxygen)make cpplint)the code is written in idiomatic C++ and changes are documented in the doxygen
the new changes are tested