FazBrowse GitHub Viewer
|
Trending
|
URL:
|
Home
Tools:
[Download Repo ZIP]
[View Raw Code]
[Original HTTPS Page]
arrayfire-binary-python-wrapper/tests/test_conv_grad.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_conv_grad.py
Copy path
More file actions
More file actions
Latest commit
History
History
History
107 lines (95 loc) · 3.5 KB
Breadcrumbs
arrayfire-binary-python-wrapper
/
tests
/
test_conv_grad.py
Copy path
File metadata and controls
107 lines (95 loc) · 3.5 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
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
import
pytest
import
arrayfire_wrapper
.
dtypes
as
dtype
import
arrayfire_wrapper
.
lib
as
wrapper
from
arrayfire_wrapper
.
lib
.
_constants
import
ConvGradient
from
tests
.
utility_functions
import
get_float_types
@
pytest
.
mark
.
parametrize
(
"grad_type"
,
[
0
,
# ConvGradient.DEFAULT
1
,
# ConvGradient.FILTER
2
,
# ConvGradient.DATA
3
,
# ConvGradient.BIAS
],
)
@
pytest
.
mark
.
parametrize
(
"dtypes"
,
get_float_types
())
def
test_convolve2_gradient_data
(
grad_type
:
int
,
dtypes
:
dtype
.
Dtype
)
->
None
:
"""Test if convolve gradient returns the correct shape with varying data type and grad type."""
incoming_gradient
=
wrapper
.
randu
((
8
,
8
),
dtypes
)
original_signal
=
wrapper
.
randu
((
10
,
10
),
dtypes
)
original_filter
=
wrapper
.
randu
((
3
,
3
),
dtypes
)
convolved_output
=
wrapper
.
randu
((
8
,
8
),
dtypes
)
strides
=
(
1
,
1
)
padding
=
(
1
,
1
)
dilation
=
(
1
,
1
)
grad_type_enum
=
ConvGradient
(
grad_type
)
result
=
wrapper
.
convolve2_gradient_nn
(
incoming_gradient
,
original_signal
,
original_filter
,
convolved_output
,
strides
,
padding
,
dilation
,
grad_type_enum
,
)
expected_shape
=
(
10
,
10
,
1
,
1
)
if
grad_type
!=
1
else
(
3
,
3
,
1
,
1
)
assert
wrapper
.
get_dims
(
result
)
==
expected_shape
,
f"Failed for grad_type:
{
grad_type_enum
}
, dtype:
{
dtypes
}
"
@
pytest
.
mark
.
parametrize
(
"invdtypes"
,
[
dtype
.
int32
,
# Integer 32-bit
dtype
.
uint32
,
# Unsigned Integer 32-bit
dtype
.
complex32
,
# Complex number with float 32-bit real and imaginary
],
)
def
test_convolve2_gradient_invalid_data
(
invdtypes
:
dtype
.
Dtype
)
->
None
:
"""Test if convolve gradient returns the correct shape with varying data type and grad type."""
with
pytest
.
raises
(
RuntimeError
):
incoming_gradient
=
wrapper
.
randu
((
8
,
8
),
invdtypes
)
original_signal
=
wrapper
.
randu
((
10
,
10
),
invdtypes
)
original_filter
=
wrapper
.
randu
((
3
,
3
),
invdtypes
)
convolved_output
=
wrapper
.
randu
((
8
,
8
),
invdtypes
)
strides
=
(
1
,
1
)
padding
=
(
1
,
1
)
dilation
=
(
1
,
1
)
grad_type_enum
=
ConvGradient
(
0
)
result
=
wrapper
.
convolve2_gradient_nn
(
incoming_gradient
,
original_signal
,
original_filter
,
convolved_output
,
strides
,
padding
,
dilation
,
grad_type_enum
,
)
expected_shape
=
(
10
,
10
,
1
,
1
)
assert
wrapper
.
get_dims
(
result
)
==
expected_shape
,
f"Failed for dtype:
{
invdtypes
}
"
# Parameterization for input shapes
@
pytest
.
mark
.
parametrize
(
"inputShape"
,
[
(
1
,
1
),
(
2
,
2
),
(
3
,
3
),
(
4
,
4
),
],
)
def
test_convolve2_gradient_input
(
inputShape
:
tuple
[
int
,
int
])
->
None
:
"""Test if convolve gradient returns the correct shape."""
incoming_gradient
=
wrapper
.
randu
((
8
,
8
),
dtype
.
f32
)
original_signal
=
wrapper
.
randu
(
inputShape
,
dtype
.
f32
)
original_filter
=
wrapper
.
randu
((
3
,
3
),
dtype
.
f32
)
convolved_output
=
wrapper
.
randu
((
8
,
8
),
dtype
.
f32
)
strides
=
(
1
,
1
)
padding
=
(
1
,
1
)
dilation
=
(
1
,
1
)
grad_type
=
ConvGradient
(
0
)
result
=
wrapper
.
convolve2_gradient_nn
(
incoming_gradient
,
original_signal
,
original_filter
,
convolved_output
,
strides
,
padding
,
dilation
,
grad_type
)
# print(array_to_string(exp, result, precision, transpose))
match
=
(
wrapper
.
get_dims
(
result
)[
0
],
wrapper
.
get_dims
(
result
)[
1
])
# print(match)
assert
inputShape
==
match
,
f"Failed for input shape:
{
inputShape
}
"
Back
|
FazBrowse Home
|
New Git URL