FazBrowse GitHub Viewer
|
Trending
|
URL:
|
Home
Tools:
[Download Repo ZIP]
[View Raw Code]
[Original HTTPS Page]
CoverageControl/python/tests/test_models.py at main · KumarRobotics/CoverageControl · GitHub
Uh oh!
There was an error while loading.
Please reload this page
.
KumarRobotics
/
CoverageControl
Public
Notifications
You must be signed in to change notification settings
Fork
6
Star
27
Code
Issues
1
Pull requests
1
Discussions
Actions
Wiki
Security and quality
0
Insights
Additional navigation options
Code
Issues
Pull requests
Discussions
Actions
Wiki
Security and quality
Insights
Expand file tree
Breadcrumbs
CoverageControl
/
python
/
tests
/
test_models.py
Copy path
More file actions
More file actions
Latest commit
History
History
History
107 lines (85 loc) · 3.76 KB
Breadcrumbs
CoverageControl
/
python
/
tests
/
test_models.py
Copy path
File metadata and controls
107 lines (85 loc) · 3.76 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
# This file is part of the CoverageControl library
#
# Author: Saurav Agarwal
# Contact: sauravag@seas.upenn.edu, agr.saurav1@gmail.com
# Repository: https://github.com/KumarRobotics/CoverageControl
#
# Copyright (c) 2024, Saurav Agarwal
#
# The CoverageControl library is free software: you can redistribute it and/or
# modify it under the terms of the GNU General Public License as published by
# the Free Software Foundation, either version 3 of the License, or (at your
# option) any later version.
#
# The CoverageControl library is distributed in the hope that it will be
# useful, but WITHOUT ANY WARRANTY; without even the implied warranty of
# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the GNU General
# Public License for more details.
#
# You should have received a copy of the GNU General Public License along with
# CoverageControl library. If not, see <https://www.gnu.org/licenses/>.
import
os
import
warnings
import
coverage_control
.
nn
as
cc_nn
import
torch
import
torch_geometric
from
coverage_control
import
IOUtils
script_dir
=
os
.
path
.
dirname
(
os
.
path
.
realpath
(
__file__
))
device
=
torch
.
device
(
"cpu"
)
with
torch
.
no_grad
():
model_file
=
os
.
path
.
join
(
script_dir
,
"data/lpac/models/model_k3_1024_state_dict.pt"
)
learning_config_file
=
os
.
path
.
join
(
script_dir
,
"data/params/learning_params.toml"
)
learning_config
=
IOUtils
.
load_toml
(
learning_config_file
)
lpac_model
=
cc_nn
.
LPAC
(
learning_config
).
to
(
device
)
lpac_model
.
load_state_dict
(
torch
.
load
(
model_file
,
weights_only
=
True
))
lpac_model
.
eval
()
use_comm_maps
=
learning_config
[
"ModelConfig"
][
"UseCommMaps"
]
map_size
=
learning_config
[
"CNNBackBone"
][
"ImageSize"
]
lpac_inputs_dict
=
torch
.
load
(
os
.
path
.
join
(
script_dir
,
"data/lpac/lpac_inputs.pt"
),
weights_only
=
True
)
lpac_inputs
=
[
torch_geometric
.
data
.
Data
.
from_dict
(
d
)
for
d
in
lpac_inputs_dict
]
def
test_cnn
():
with
torch
.
no_grad
():
ref_cnn_outputs
=
torch
.
load
(
os
.
path
.
join
(
script_dir
,
"data/lpac/cnn_outputs.pt"
),
weights_only
=
True
)
cnn_model
=
lpac_model
.
cnn_backbone
.
to
(
device
).
eval
()
for
i
in
range
(
0
,
len
(
lpac_inputs
)):
cnn_output
=
cnn_model
(
lpac_inputs
[
i
].
x
)
is_close
=
torch
.
allclose
(
cnn_output
,
ref_cnn_outputs
[
i
],
atol
=
1e-4
)
if
not
is_close
:
error
=
torch
.
sum
(
torch
.
abs
(
cnn_output
-
ref_cnn_outputs
[
i
]))
print
(
f"Error:
{
error
}
at
{
i
}
"
)
assert
is_close
break
is_equal
=
torch
.
equal
(
cnn_output
,
ref_cnn_outputs
[
i
])
if
not
is_equal
and
is_close
:
error
=
torch
.
sum
(
torch
.
abs
(
cnn_output
-
ref_cnn_outputs
[
i
]))
print
(
f"Error:
{
error
}
at
{
i
}
"
)
warnings
.
warn
(
"Outputs are close but not equal"
)
def
test_lpac
():
with
torch
.
no_grad
():
ref_lpac_outputs
=
torch
.
load
(
os
.
path
.
join
(
script_dir
,
"data/lpac/lpac_outputs.pt"
),
weights_only
=
True
)
for
i
in
range
(
0
,
len
(
lpac_inputs
)):
lpac_output
=
lpac_model
(
lpac_inputs
[
i
])
is_close
=
torch
.
allclose
(
lpac_output
,
ref_lpac_outputs
[
i
],
atol
=
1e-4
)
if
not
is_close
:
error
=
torch
.
sum
(
torch
.
abs
(
lpac_output
-
ref_lpac_outputs
[
i
]))
print
(
f"Error:
{
error
}
at
{
i
}
"
)
assert
is_close
break
is_equal
=
torch
.
equal
(
lpac_output
,
ref_lpac_outputs
[
i
])
if
not
is_equal
and
is_close
:
error
=
torch
.
sum
(
torch
.
abs
(
lpac_output
-
ref_lpac_outputs
[
i
]))
print
(
f"Error:
{
error
}
at
{
i
}
"
)
warnings
.
warn
(
"Outputs are close but not equal"
)
if
__name__
==
"__main__"
:
test_cnn
()
test_lpac
()
Back
|
FazBrowse Home
|
New Git URL