FazBrowse GitHub Viewer
|
Trending
|
URL:
|
Home
Tools:
[Download Repo ZIP]
[View Raw Code]
[Original HTTPS Page]
arrayfire-binary-python-wrapper/tests/test_diag.py at master · arrayfire/arrayfire-binary-python-wrapper · GitHub
arrayfire
arrayfire-binary-python-wrapper
Repository navigation
Code
Issues
1
(1)
Pull requests
2
(2)
Discussions
Actions
Projects
Security and quality
Insights
Expand file tree
Breadcrumbs
arrayfire-binary-python-wrapper
/
tests
/
test_diag.py
Copy path
More file actions
More file actions
Latest commit
History
History
History
43 lines (31 loc) · 1.53 KB
Breadcrumbs
arrayfire-binary-python-wrapper
/
tests
/
test_diag.py
Copy path
File metadata and controls
43 lines (31 loc) · 1.53 KB
Raw
Copy raw file
Download raw file
Open symbols panel
Edit and raw actions
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
import
pytest
import
arrayfire_wrapper
.
dtypes
as
dtypes
import
arrayfire_wrapper
.
lib
as
wrapper
@
pytest
.
mark
.
parametrize
(
"diagonal_shape"
, [(
2
,), (
10
,), (
100
,), (
1000
,)])
def
test_diagonal_shape
(
diagonal_shape
:
tuple
)
->
None
:
"""Test if diagonal array is keeping the shape of the passed into the input array"""
in_arr
=
wrapper
.
constant
(
1
,
diagonal_shape
,
dtypes
.
s16
)
diag_array
=
wrapper
.
diag_create
(
in_arr
,
0
)
extracted_diagonal
=
wrapper
.
diag_extract
(
diag_array
,
0
)
assert
wrapper
.
get_dims
(
extracted_diagonal
)[
0
:
len
(
diagonal_shape
)]
==
diagonal_shape
# noqa: E203
@
pytest
.
mark
.
parametrize
(
"diagonal_shape"
, [(
2
,), (
10
,), (
100
,), (
1000
,)])
def
test_diagonal_val
(
diagonal_shape
:
tuple
)
->
None
:
"""Test if diagonal array is keeping the same value as that of the values passed into the input array"""
dtype
=
dtypes
.
s16
in_arr
=
wrapper
.
constant
(
1
,
diagonal_shape
,
dtype
)
diag_array
=
wrapper
.
diag_create
(
in_arr
,
0
)
extracted_diagonal
=
wrapper
.
diag_extract
(
diag_array
,
0
)
assert
wrapper
.
get_scalar
(
extracted_diagonal
,
dtype
)
==
wrapper
.
get_scalar
(
in_arr
,
dtype
)
@
pytest
.
mark
.
parametrize
(
"diagonal_shape"
,
[
(
10
,
10
,
10
),
(
100
,
100
,
100
,
100
),
],
)
def
test_invalid_diagonal
(
diagonal_shape
:
tuple
)
->
None
:
"""Test if an invalid diagonal shape is being properly handled"""
with
pytest
.
raises
(
RuntimeError
):
in_arr
=
wrapper
.
constant
(
1
,
diagonal_shape
,
dtypes
.
s16
)
diag_array
=
wrapper
.
diag_create
(
in_arr
,
0
)
wrapper
.
diag_extract
(
diag_array
,
0
)
Back
|
FazBrowse Home
|
New Git URL