From ec7bba906eeb11cef587f7aba181c1337232f709 Mon Sep 17 00:00:00 2001 From: Arihan Yadav Date: Sat, 22 Aug 2026 21:13:52 -0700 Subject: [PATCH] Fix boolean scalar support in portable arange --- kernels/portable/cpu/op_arange.cpp | 20 ++++++++++++++++---- kernels/test/op_arange_test.cpp | 20 ++++++++++++++++++++ 2 files changed, 36 insertions(+), 4 deletions(-) diff --git a/kernels/portable/cpu/op_arange.cpp b/kernels/portable/cpu/op_arange.cpp index 0013f1a4d0f..25f214152df 100644 --- a/kernels/portable/cpu/op_arange.cpp +++ b/kernels/portable/cpu/op_arange.cpp @@ -20,10 +20,22 @@ namespace torch { namespace executor { namespace native { +namespace { + +bool extract_arange_scalar(const Scalar& scalar, double* out) { + if (scalar.isBoolean()) { + *out = static_cast(scalar.to()); + return true; + } + return utils::extract_scalar(scalar, out); +} + +} // namespace + Tensor& arange_out(KernelRuntimeContext& ctx, const Scalar& end, Tensor& out) { double end_val = 0; ET_KERNEL_CHECK( - ctx, utils::extract_scalar(end, &end_val), InvalidArgument, out); + ctx, extract_arange_scalar(end, &end_val), InvalidArgument, out); ET_KERNEL_CHECK( ctx, check_arange_args(0.0, end_val, 1.0, out), InvalidArgument, out); @@ -53,15 +65,15 @@ Tensor& arange_start_out( double d_start = 0; ET_KERNEL_CHECK( - ctx, utils::extract_scalar(start, &d_start), InvalidArgument, out); + ctx, extract_arange_scalar(start, &d_start), InvalidArgument, out); double d_end = 0; ET_KERNEL_CHECK( - ctx, utils::extract_scalar(end, &d_end), InvalidArgument, out); + ctx, extract_arange_scalar(end, &d_end), InvalidArgument, out); double d_step = 0; ET_KERNEL_CHECK( - ctx, utils::extract_scalar(step, &d_step), InvalidArgument, out); + ctx, extract_arange_scalar(step, &d_step), InvalidArgument, out); ET_KERNEL_CHECK( ctx, diff --git a/kernels/test/op_arange_test.cpp b/kernels/test/op_arange_test.cpp index 0a129ae8046..d63c15d41bd 100644 --- a/kernels/test/op_arange_test.cpp +++ b/kernels/test/op_arange_test.cpp @@ -113,6 +113,15 @@ TEST_F(OpArangeOutTest, FloatNumberNotEqualIntSupport) { EXPECT_TENSOR_EQ(out, expected); } +TEST_F(OpArangeOutTest, BooleanEndSupported) { + TensorFactory tf; + + Tensor out = tf.zeros({1}); + Tensor expected = tf.make({1}, {0}); + + EXPECT_TENSOR_EQ(op_arange_out(Scalar(true), out), expected); +} + TEST_F(OpArangeOutTest, OutDimUnsupportedDie) { ET_SKIP_IF( torch::executor::testing::SupportedFeatures::get()->is_aten, @@ -195,6 +204,17 @@ TEST_F(OpArangeStartOutTest, FloatNumberNotEqualIntSupport) { EXPECT_TENSOR_EQ(out, expected); } +TEST_F(OpArangeStartOutTest, BooleanStartAndStepSupported) { + TensorFactory tf; + + Tensor out = tf.zeros({3}); + Tensor expected = tf.make({3}, {0, 1, 2}); + + EXPECT_TENSOR_EQ( + op_arange_start_out(Scalar(false), Scalar(3), Scalar(true), out), + expected); +} + TEST_F(OpArangeStartOutTest, OutDimUnsupportedDie) { ET_SKIP_IF( torch::executor::testing::SupportedFeatures::get()->is_aten,