-
Notifications
You must be signed in to change notification settings - Fork 493
Expand file tree
/
Copy pathtest_dsplit.py
More file actions
65 lines (56 loc) · 2.07 KB
/
Copy pathtest_dsplit.py
File metadata and controls
65 lines (56 loc) · 2.07 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
# Copyright 2026 FlagOS Contributors
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import pytest
import torch
import flag_gems
from . import accuracy_utils as utils
DSPLIT_CONFIGS = [
# (shape, indices_or_sections)
# Integer splits (equal chunks)
((4, 6, 8), 2),
((8, 4, 12), 4),
((12, 8, 16), 4),
((16, 16, 8), 2),
((20, 10, 15), 5),
# List splits (custom indices)
((4, 6, 8), [2]),
((4, 6, 8), [4]),
((6, 8, 10), [2, 6]),
((8, 4, 12), [3, 6, 9]),
((10, 5, 15), [5, 10]),
# 4D tensors
((4, 6, 8, 3), 2),
((8, 4, 12, 5), 4),
((6, 8, 10, 7), [2, 6]),
((10, 5, 15, 3), [5, 10]),
]
@pytest.mark.dsplit
@pytest.mark.parametrize("shape, indices_or_sections", DSPLIT_CONFIGS)
def test_accuracy_dsplit(shape, indices_or_sections):
inp = torch.randn(shape, dtype=torch.float32, device=flag_gems.device)
ref_inp = utils.to_reference(inp, True)
if isinstance(indices_or_sections, int):
ref_out = torch.ops.aten.dsplit.int(ref_inp, indices_or_sections)
else:
ref_out = torch.ops.aten.dsplit.array(ref_inp, indices_or_sections)
with flag_gems.use_gems():
if isinstance(indices_or_sections, int):
res_out = torch.ops.aten.dsplit.int(inp, indices_or_sections)
else:
res_out = torch.ops.aten.dsplit.array(inp, indices_or_sections)
assert len(res_out) == len(
ref_out
), f"Length mismatch: {len(res_out)} vs {len(ref_out)}"
for i, (res_chunk, ref_chunk) in enumerate(zip(res_out, ref_out)):
utils.gems_assert_close(res_chunk, ref_chunk, torch.float32)