@@ -82,6 +82,17 @@ struct type_map<torch::executor::ArrayRef<T>> final {
8282 using type = at::ArrayRef<typename type_map<T>::type>;
8383};
8484
85+ // Tuple.
86+ template <class ... Args>
87+ struct type_map <std::tuple<Args...>> final {
88+ using type = std::tuple<typename type_map<Args>::type...>;
89+ };
90+
91+ template <class ... Args>
92+ struct type_map <std::tuple<Args...>&> final {
93+ using type = std::tuple<typename type_map<Args>::type...>&;
94+ };
95+
8596template <typename T>
8697struct remove_const_ref final {
8798 using type = std::remove_const_t <std::remove_reference_t <T>>;
@@ -232,6 +243,77 @@ struct type_convert<torch::executor::ArrayRef<F>, c10::ArrayRef<T>> final {
232243 }
233244};
234245
246+ // Tuples: birdirectional.
247+ template <class ... F, class ... T>
248+ struct type_convert <std::tuple<F...>, std::tuple<T...>> final {
249+ public:
250+ // We need to remove references from the output type because converters expect
251+ // to return by value. For example, the tensor converter creates a temporary
252+ // tensor and returns it by value - so we can't return an lvalue reference to
253+ // it.
254+ using InType = std::tuple<F...>;
255+ using OutType = std::tuple<std::remove_reference_t <T>...>;
256+
257+ std::tuple<type_convert<F, std::remove_reference_t <T>>...> converters;
258+
259+ explicit type_convert (std::tuple<F...> value)
260+ : converters(make_converters(value, std::index_sequence_for<F...>{})) {}
261+
262+ template <size_t ... Is>
263+ static decltype (converters) make_converters(
264+ std::tuple<F...> value,
265+ std::index_sequence<Is...>) {
266+ return std::make_tuple (
267+ // For each element in the input/output tuple, create a type_convert
268+ // struct.
269+ type_convert<
270+ // Get the type of Ith element of the input and output.
271+ std::tuple_element_t <Is, InType>,
272+ std::tuple_element_t <Is, OutType>
273+ // Instantiate the converter with the Ith element of the input.
274+ >(std::get<Is>(value))...);
275+ }
276+
277+ template <size_t ... Is>
278+ OutType call_impl (std::index_sequence<Is...>) {
279+ // Invoke the converter on each element of the tuple.
280+ return std::make_tuple (std::get<Is>(converters).call ()...);
281+ }
282+ OutType call () {
283+ return call_impl (std::make_index_sequence<sizeof ...(F)>{});
284+ }
285+ };
286+
287+ template <class T >
288+ struct is_tuple : std::false_type {};
289+
290+ template <class ... Args>
291+ struct is_tuple <std::tuple<Args...>> : std::true_type {};
292+
293+ // A utility struct to extract the out arguments from a function.
294+ // Returns a tuple of Args[N...], where N is the index of the first out
295+ // argument.
296+ template <size_t N, typename ... Args>
297+ struct extract_out_args ;
298+
299+ template <size_t N, typename ... Args>
300+ struct extract_out_args {
301+ static_assert (sizeof ...(Args) >= N, " Output index out of range." );
302+ static constexpr size_t n_out = sizeof ...(Args) - N;
303+ using tuple_type = std::tuple<Args...>;
304+
305+ template <size_t ... I>
306+ static auto get_impl (tuple_type&& t, std::index_sequence<I...>) {
307+ return std::forward_as_tuple (
308+ std::get<N + I>(std::forward<tuple_type>(t))...);
309+ }
310+
311+ static auto get (tuple_type&& t) {
312+ return get_impl (
313+ std::forward<tuple_type>(t), std::make_index_sequence<n_out>{});
314+ }
315+ };
316+
235317template <class F , F f, typename N = int , N index = N(-1 )>
236318struct wrapper_impl ;
237319
@@ -244,17 +326,23 @@ struct wrapper_impl<R (*)(Args...), f, int, N> {
244326 using TupleConvertsType =
245327 std::tuple<type_convert<typename type_map<Args>::type, Args>...>;
246328 using TupleArgsType = std::tuple<typename type_map<Args>::type...>;
329+
247330 static constexpr size_t num_args = sizeof ...(Args);
331+ static constexpr size_t num_out_args = sizeof ...(Args) - N;
332+ static constexpr bool is_output_tuple = is_tuple<ReturnType>::value;
248333 static_assert (
249- (N < num_args &&
250- std::is_same_v<
251- executorch::extension::kernel_util_internal::element_t <
252- N,
253- executorch::extension::kernel_util_internal::typelist<Args...>>,
254- R>) ||
334+ N < num_args,
335+ " The index of the out tensor can't be greater or equal to num_args." );
336+ static_assert (
337+ is_output_tuple ||
338+ std::is_same_v<
339+ executorch::extension::kernel_util_internal::element_t <
340+ N,
341+ executorch::extension::kernel_util_internal::typelist<
342+ Args...>>,
343+ R> ||
255344 N == -1 ,
256- " The index of the out tensor can't be greater or equal to num_args and "
257- " the Nth argument type has to be the same as the return type." );
345+ " The Nth argument type has to be the same as the return type." );
258346
259347 static ReturnType wrap (typename type_map<Args>::type... args) {
260348 // The wrapped function that takes ATen argument types, convert them into
@@ -265,20 +353,26 @@ struct wrapper_impl<R (*)(Args...), f, int, N> {
265353 type_convert<typename type_map<Args>::type, Args>(args)...);
266354 R result =
267355 call_functor_with_args (converts, std::make_index_sequence<num_args>());
268- typename std::remove_reference<ReturnType>::type converted_result =
269- type_convert<R, ReturnType>(result).call ();
356+ auto converted_result = type_convert<R, ReturnType>(result).call ();
270357 if constexpr (N == -1 ) {
271358 return converted_result;
272- } else {
359+ } else if constexpr (is_output_tuple) {
360+ auto out_args =
361+ extract_out_args<N, typename type_map<Args>::type...>::get (
362+ std::move (args_tuple));
363+
364+ return resize_and_copy_outputs (
365+ std::move (converted_result),
366+ std::move (out_args),
367+ std::make_index_sequence<num_out_args>{});
368+ } else { // Non-tuple return type.
273369 static_assert (
274370 std::is_same_v<
275371 typename std::remove_reference<ReturnType>::type,
276372 at::Tensor>,
277373 " Only support at::Tensor-like return" );
278374 ReturnType out = std::get<N>(args_tuple);
279- at::native::resize_output (out, converted_result.sizes ());
280- out.copy_ (converted_result);
281- return out;
375+ return resize_and_copy_output (converted_result, out);
282376 }
283377 }
284378
@@ -289,6 +383,25 @@ struct wrapper_impl<R (*)(Args...), f, int, N> {
289383 std::index_sequence<indices...>) {
290384 return f (std::get<indices>(converts).call ()...);
291385 }
386+
387+ template <class A >
388+ static at::Tensor& resize_and_copy_output (
389+ A converted_result,
390+ at::Tensor& out) {
391+ at::native::resize_output (out, converted_result.sizes ());
392+ out.copy_ (converted_result);
393+ return out;
394+ }
395+
396+ template <class A , class B , size_t ... Is>
397+ static ReturnType resize_and_copy_outputs (
398+ A&& converted_result,
399+ B&& out,
400+ std::index_sequence<Is...>) {
401+ return std::forward_as_tuple (resize_and_copy_output (
402+ std::get<Is>(std::forward<A>(converted_result)),
403+ std::get<Is>(std::forward<B>(out)))...);
404+ }
292405};
293406
294407} // namespace internal
0 commit comments