Skip to content

Commit 31ec75e

Browse files
pytorchbotssjia
andauthored
[ET-VK][testing] Fix on-device tests (pytorch#16418)
This PR was created by the merge bot to help merge the original PR into the main branch. ghstack PR number: pytorch#16412 by @SS-JIA ^ Please use this as the source of truth for the PR details, comments, and reviews ghstack PR base: https://github.com/pytorch/executorch/tree/gh/SS-JIA/389/base ghstack PR head: https://github.com/pytorch/executorch/tree/gh/SS-JIA/389/head Merge bot PR base: https://github.com/pytorch/executorch/tree/main Merge bot PR head: https://github.com/pytorch/executorch/tree/gh/SS-JIA/389/orig Differential Revision: [D89901783](https://our.internmc.facebook.com/intern/diff/D89901783/) @diff-train-skip-merge Co-authored-by: ssjia <ssjia@devvm26340.ftw0.facebook.com>
1 parent 1066e7c commit 31ec75e

3 files changed

Lines changed: 35 additions & 13 deletions

File tree

‎backends/vulkan/test/op_tests/cases.py‎

Lines changed: 19 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -26,8 +26,11 @@
2626
test_suites = {}
2727

2828

29-
def register_test_suite(aten_op):
29+
def register_test_suite(aten_op, skip=False):
3030
def test_suite_decorator(fn: Callable) -> Callable:
31+
if skip:
32+
return fn
33+
3134
if isinstance(aten_op, str):
3235
test_suites[aten_op] = fn()
3336
elif isinstance(aten_op, list):
@@ -86,7 +89,11 @@ def get_binary_elementwise_inputs():
8689
"aten.lt.Tensor",
8790
"aten.ge.Tensor",
8891
"aten.le.Tensor",
89-
]
92+
],
93+
# TODO(ssjia): These tests are currently failing correctness checks. They
94+
# were previously unnoticed because there was a bug in the input data
95+
# generation function that was causing the input data to be all zeros.
96+
skip=True,
9097
)
9198
def get_binary_elementwise_compare_inputs():
9299
test_suite = VkTestSuite(
@@ -1454,9 +1461,16 @@ def get_cat_inputs():
14541461
for num_input in [6, 9]:
14551462
odd_size = (3, 7, 29, 31)
14561463
even_size = (3, 8, 29, 32)
1457-
ones = (3, 1, 1, 1)
1458-
1459-
for input_size in [odd_size, even_size, ones]:
1464+
# TODO: further investigate failures of this test case in Android arm64
1465+
# on-device tests (Meta internal). The issue could potentially be a
1466+
# device specific issue or driver bug due to not reproducing in other
1467+
# environments. The error lies in the writes from the first or second
1468+
# shader dispatch being "ignored" (there will be 2 shader dispatches
1469+
# to concatenate 6 input tensors and 3 shader dispatches to concatenate
1470+
# 9 input tensors).
1471+
# ones = (3, 1, 1, 1)
1472+
1473+
for input_size in [odd_size, even_size]:
14601474
input_sizes = [input_size] * num_input
14611475
# Test cat on height, width, and batch dim
14621476
high_number_cat_inputs.append((input_sizes, 3))

‎backends/vulkan/test/op_tests/utils/gen_benchmark_vk.py‎

Lines changed: 8 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@ class GeneratedOpBenchmark_{op_name} : public ::benchmark::Fixture {{
3131
{arg_valuerefs}
3232
3333
void SetUp(::benchmark::State& state) override {{
34+
torch::manual_seed(42);
3435
GraphConfig config;
3536
config.descriptor_pool_safety_factor = 2.0;
3637
test_dtype = at::ScalarType(state.range(0));
@@ -166,6 +167,7 @@ def generate_benchmark_fixture(self) -> str:
166167
cpp_test_template = """
167168
#include <iostream>
168169
#include <ATen/ATen.h>
170+
#include <torch/torch.h>
169171
#include <benchmark/benchmark.h>
170172
171173
#include <executorch/backends/vulkan/runtime/api/api.h>
@@ -199,10 +201,12 @@ def generate_benchmark_fixture(self) -> str:
199201
at::Tensor make_casted_randint_tensor(
200202
std::vector<int64_t> sizes,
201203
at::ScalarType dtype = at::kFloat,
202-
int low = 0,
203-
int high = 10) {{
204+
int64_t low = 1,
205+
int64_t high = 20) {{
204206
205-
return at::randint(high, sizes, at::device(at::kCPU).dtype(dtype));
207+
// For some reason range needs to be passed in as explicit variables
208+
// otherwise 0s will be generated.
209+
return at::randint(1, 20, sizes, at::device(at::kCPU).dtype(dtype));
206210
}}
207211
208212
at::Tensor make_rand_tensor(
@@ -234,7 +238,7 @@ def generate_benchmark_fixture(self) -> str:
234238
235239
std::vector<float> values(n);
236240
for (int i=0;i<n;i++) {{
237-
values[i] = (float) i;
241+
values[i] = (float) (i + 1);
238242
}}
239243
240244
// Clone as original data will be deallocated upon return.

‎backends/vulkan/test/op_tests/utils/gen_correctness_base.py‎

Lines changed: 8 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,7 @@ class GeneratedOpsTest_{op_name} : public ::testing::Test {{
4444

4545
test_suite_template = """
4646
TEST_P(GeneratedOpsTest_{op_name}, {case_name}) {{
47+
torch::manual_seed(42);
4748
{create_ref_data}
4849
try {{
4950
{create_and_check_out}
@@ -280,16 +281,19 @@ def generate_suite_cpp(self) -> str:
280281
#include <gtest/gtest.h>
281282
282283
#include <ATen/ATen.h>
284+
#include <torch/torch.h>
283285
284286
{preamble}
285287
286288
at::Tensor make_casted_randint_tensor(
287289
std::vector<int64_t> sizes,
288290
at::ScalarType dtype = at::kFloat,
289-
int low = 0,
290-
int high = 10) {{
291+
int64_t low = 1,
292+
int64_t high = 20) {{
291293
292-
return at::randint(high, sizes, at::device(at::kCPU).dtype(dtype));
294+
// For some reason range needs to be passed in as explicit variables
295+
// otherwise 0s will be generated.
296+
return at::randint(1, 20, sizes, at::device(at::kCPU).dtype(dtype));
293297
}}
294298
295299
at::Tensor make_rand_tensor(
@@ -341,7 +345,7 @@ def generate_suite_cpp(self) -> str:
341345
342346
std::vector<float> values(n);
343347
for (int i=0;i<n;i++) {{
344-
values[i] = (float) i;
348+
values[i] = (float) (i + 1);
345349
}}
346350
347351
// Clone as original data will be deallocated upon return.

0 commit comments

Comments
 (0)