Skip to content

Commit 74ea584

Browse files
authored
Support tuple outputs in aten bridge (pytorch#16848)
Summary: Update the aten bridge logic in make_aten_functor_from_et... to handle kernels that return tuples of tensors. This motivating use case is the box_with_nms_limit custom op, which returns a tuple of 6 output tensors. Differential Revision: D91375148
1 parent 61ae2f6 commit 74ea584

2 files changed

Lines changed: 219 additions & 14 deletions

File tree

‎extension/aten_util/make_aten_functor_from_et_functor.h‎

Lines changed: 127 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -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+
8596
template <typename T>
8697
struct 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+
235317
template <class F, F f, typename N = int, N index = N(-1)>
236318
struct 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

‎extension/aten_util/test/make_aten_functor_from_et_functor_test.cpp‎

Lines changed: 92 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -90,6 +90,20 @@ Tensor& sum_arrayref_optional_tensor_out(
9090
return out;
9191
}
9292

93+
std::tuple<Tensor&, Tensor&>
94+
add_constant_tuple_out(const Tensor& a, Tensor& out1, Tensor& out2) {
95+
auto a_data = a.const_data_ptr<int32_t>();
96+
auto out1_data = out1.mutable_data_ptr<int32_t>();
97+
auto out2_data = out2.mutable_data_ptr<int32_t>();
98+
99+
for (int i = 0; i < a.size(0); i++) {
100+
out1_data[i] = a_data[i] + 1;
101+
out2_data[i] = a_data[i] + 2;
102+
}
103+
104+
return {out1, out2};
105+
}
106+
93107
Tensor& quantized_embedding_byte_out(
94108
const Tensor& weight,
95109
const Tensor& weight_scales,
@@ -134,6 +148,24 @@ TEST_F(MakeATenFunctorFromETFunctorTest, TestTypeMap_Tensor) {
134148
const at::Tensor&>::value));
135149
}
136150

151+
TEST_F(MakeATenFunctorFromETFunctorTest, TestTypeMap_Tuple_Tensor2x) {
152+
EXPECT_TRUE(
153+
(std::is_same<
154+
type_map<
155+
std::tuple<torch::executor::Tensor, torch::executor::Tensor>>::
156+
type,
157+
std::tuple<at::Tensor, at::Tensor>>::value));
158+
}
159+
160+
TEST_F(MakeATenFunctorFromETFunctorTest, TestTypeMap_Tuple_TensorRef3x) {
161+
EXPECT_TRUE((std::is_same<
162+
type_map<std::tuple<
163+
torch::executor::Tensor&,
164+
torch::executor::Tensor&,
165+
torch::executor::Tensor&>>::type,
166+
std::tuple<at::Tensor&, at::Tensor&, at::Tensor&>>::value));
167+
}
168+
137169
TEST_F(MakeATenFunctorFromETFunctorTest, TestTypeMap_Optionals) {
138170
// Scalar.
139171
EXPECT_TRUE((std::is_same<
@@ -189,6 +221,34 @@ TEST_F(MakeATenFunctorFromETFunctorTest, TestConvert_Tensor) {
189221
EXPECT_TRUE((std::is_same<decltype(at_out), at::Tensor>::value));
190222
}
191223

224+
TEST_F(MakeATenFunctorFromETFunctorTest, TestConvert_TupleTensor) {
225+
// Convert at to et.
226+
at::Tensor at_in1 = torch::tensor({1});
227+
at::Tensor at_in2 = torch::tensor({2});
228+
auto et = type_convert<
229+
std::tuple<at::Tensor, at::Tensor>,
230+
std::tuple<torch::executor::Tensor, torch::executor::Tensor>>(
231+
std::make_tuple(at_in1, at_in2))
232+
.call();
233+
EXPECT_TRUE((std::is_same<
234+
decltype(et),
235+
std::tuple<torch::executor::Tensor, torch::executor::Tensor>>::
236+
value));
237+
238+
// Convert et to at.
239+
torch::executor::testing::TensorFactory<ScalarType::Int> tf;
240+
torch::executor::Tensor et_in1 = tf.ones({3});
241+
torch::executor::Tensor et_in2 = tf.ones({4});
242+
auto at_out =
243+
type_convert<
244+
std::tuple<torch::executor::Tensor, torch::executor::Tensor>,
245+
std::tuple<at::Tensor, at::Tensor>>(std::make_tuple(et_in1, et_in2))
246+
.call();
247+
EXPECT_TRUE(
248+
(std::is_same<decltype(at_out), std::tuple<at::Tensor, at::Tensor>>::
249+
value));
250+
}
251+
192252
TEST_F(MakeATenFunctorFromETFunctorTest, TestConvert_OptionalScalar) {
193253
// Convert optional at to et.
194254
auto optional_at_in = std::optional<int64_t>();
@@ -305,6 +365,9 @@ TORCH_LIBRARY(my_op, m) {
305365
m.def("add_optional_tensor.out", WRAP_TO_ATEN(add_optional_tensor_out, 2));
306366
m.def("sum_arrayref_scalar.out", WRAP_TO_ATEN(sum_arrayref_scalar_out, 1));
307367
m.def("sum_arrayref_tensor.out", WRAP_TO_ATEN(sum_arrayref_tensor_out, 1));
368+
m.def(
369+
"add_constant_tuple.out(Tensor a, *, Tensor(a!) out1, Tensor(b!) out2) -> (Tensor(a!), Tensor(b!))",
370+
WRAP_TO_ATEN(add_constant_tuple_out, 1));
308371
m.def(
309372
"sum_arrayref_optional_tensor.out",
310373
WRAP_TO_ATEN(sum_arrayref_optional_tensor_out, 1));
@@ -422,6 +485,35 @@ TEST_F(MakeATenFunctorFromETFunctorTest, TestWrap_ArrayRefOptional) {
422485
EXPECT_EQ(stack[0].toTensor().const_data_ptr<int64_t>()[0], 4);
423486
}
424487

488+
TEST_F(MakeATenFunctorFromETFunctorTest, TestWrap_TupleOut) {
489+
at::Tensor a =
490+
torch::tensor({1, 2, 3}, torch::TensorOptions().dtype(torch::kInt32));
491+
at::Tensor out1 =
492+
torch::tensor({0, 0, 0}, torch::TensorOptions().dtype(torch::kInt32));
493+
at::Tensor out2 =
494+
torch::tensor({0, 0, 0}, torch::TensorOptions().dtype(torch::kInt32));
495+
496+
auto op = c10::Dispatcher::singleton().findSchema(
497+
{"my_op::add_constant_tuple", "out"});
498+
EXPECT_TRUE(op.has_value());
499+
torch::jit::Stack stack = {a, out1, out2};
500+
op.value().callBoxed(&stack);
501+
502+
EXPECT_EQ(stack.size(), 2);
503+
504+
// Verify out1 contains a + 1
505+
auto result1 = stack[0].toTensor();
506+
EXPECT_EQ(result1.const_data_ptr<int32_t>()[0], 2);
507+
EXPECT_EQ(result1.const_data_ptr<int32_t>()[1], 3);
508+
EXPECT_EQ(result1.const_data_ptr<int32_t>()[2], 4);
509+
510+
// Verify out2 contains a + 2
511+
auto result2 = stack[1].toTensor();
512+
EXPECT_EQ(result2.const_data_ptr<int32_t>()[0], 3);
513+
EXPECT_EQ(result2.const_data_ptr<int32_t>()[1], 4);
514+
EXPECT_EQ(result2.const_data_ptr<int32_t>()[2], 5);
515+
}
516+
425517
TEST_F(MakeATenFunctorFromETFunctorTest, TestConvert_ConstRefOptionals) {
426518
// Test const optional scalar conversion
427519
const std::optional<int64_t> const_optional_at_in =

0 commit comments

Comments
 (0)