-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathbenchmarkmain.py
More file actions
88 lines (72 loc) · 2.43 KB
/
Copy pathbenchmarkmain.py
File metadata and controls
88 lines (72 loc) · 2.43 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
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
import argparse
import os
import pybullet as p
from ball.delayed_pybullet_ball import DelayedPybulletBall
from benchmarks.stpos_benchmark import StPosBenchmark
from benchmarks.wind_benchmark import WindBenchmark
from position_prediction.DES import DESPredicter
from benchmarks.force_benchmark import ForceBenchmark
from position_prediction.no_prediction import NoPredictionPredicter
from trackers.concurrent_ball_tracker import ConcurrentPredictingBallTracker
from utils.environment import init_env_and_load_assets
parser = argparse.ArgumentParser()
DEFAULT_FILE_NAME = "prediction_error"
DEFAULT_N_PREDICT = 1
DEFAULT_N_DELAYED = 0
DEFAULT_FETCH_TIME = 1 / 30 # 20 is the maximum number of camera outputs per second.
parser.add_argument("--file_name", default=DEFAULT_FILE_NAME, type=str)
parser.add_argument(
"-d", type=int, help="N_DELAYED parameter", default=DEFAULT_N_DELAYED
)
parser.add_argument(
"--benchmark-type", type=str, help="Specify the benchmark type.", default="FORCE"
)
parser.add_argument(
"-p", type=int, help="N_PREDICT parameter", default=DEFAULT_N_PREDICT
)
parser.add_argument(
"-f", type=float, help="FETCH_TIME parameter", default=DEFAULT_FETCH_TIME
)
parser.add_argument(
"--no-prediction", help="USE no_prediciton predicter", action="store_true"
)
parser.add_argument("--delete", action="store_true")
args = parser.parse_args()
N_DELAYED = args.d
N_PREDICT = args.p
FETCH_TIME = args.f
NO_PREDICTION = args.no_prediction
BENCHMARK_TYPE = args.benchmark_type
(
ball_controller,
ball,
paddle,
wind_controllers,
force_controllers,
) = init_env_and_load_assets(p)
paddle.create_joint_controllers()
if NO_PREDICTION:
predicter = NoPredictionPredicter()
else:
predicter = DESPredicter(1)
tracker = ConcurrentPredictingBallTracker(
DelayedPybulletBall(ball, N_DELAYED), paddle, predicter, FETCH_TIME
)
file_name = args.file_name
csv_name = file_name + ".csv"
plot_name = file_name + ".png"
if BENCHMARK_TYPE == "POSITION":
benchmark = StPosBenchmark(tracker, ball)
elif BENCHMARK_TYPE == "WIND":
benchmark = WindBenchmark(tracker, ball)
else:
benchmark = ForceBenchmark(tracker, ball, [(0.225, 0.225), (0.1, 0.1), (-0.2, 0.1)])
benchmark.run_benchmark(p, paddle, 3, (400.0, 1.0, 70.0))
print("SAVING PLOTS")
if BENCHMARK_TYPE == "FORCE":
benchmark.get_prediction_error(plot_name)
benchmark.plot()
input("CONTINUE?")
if args.delete:
os.remove(csv_name)
os.remove(plot_name)