-
Notifications
You must be signed in to change notification settings - Fork 3
Expand file tree
/
Copy pathDockerfile
More file actions
128 lines (110 loc) · 4.83 KB
/
Copy pathDockerfile
File metadata and controls
128 lines (110 loc) · 4.83 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
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
FROM repo.irsl.eiiris.tut.ac.jp/irsl_system:one
ARG TORCH_VER=2.9
###
RUN (cd /; git clone https://github.com/IRSL-tut/RoboManipBaselines.git --recursive)
WORKDIR /RoboManipBaselines
RUN apt update -q -qq && \
apt install -q -qq -y ffmpeg python3-venv libnppicc12 && \
apt clean && \
rm -rf /var/lib/apt/lists/
# RUN python3 -m venv /irsl_venv --copies --system-site-packages
RUN python3 -m venv /irsl_venv --copies
## install pytorch
RUN <<EOF
if [ -e /irsl_venv/bin/activate ]; then
source /irsl_venv/bin/activate
fi
mkdir -p /opt/python
pip install --target /opt/python iceoryx2==0.7.0
#
if [ ${TORCH_VER} == '2.9' ]; then
pip install --break-system-packages torch==2.9.0 torchvision torchcodec==0.8
elif [ ${TORCH_VER} == '2.8' ]; then
pip install --break-system-packages torch==2.8.0 torchvision==0.23.0 torchaudio==2.8.0 torchcodec==0.6 --index-url https://download.pytorch.org/whl/cu128
elif [ ${TORCH_VER} == '2.7' ]; then
pip install --break-system-packages torch==2.7.1 torchvision==0.22.1 torchaudio==2.7.1 torchcodec==0.5 --index-url https://download.pytorch.org/whl/cu126
else
set -e
[ 0 -eq 1 ] ## failed
fi
EOF
RUN source /irsl_venv/bin/activate && \
cd /RoboManipBaselines && \
pip install -e .[act] && \
cd third_party/act/detr && \
pip install -e .
## add for fix torch version
# sed -i -e 's@"torch"@"torch<2.9"@' pyproject.toml && \
## SARNN
RUN source /irsl_venv/bin/activate && \
cd /RoboManipBaselines && \
pip install -e .[sarnn] && \
cd third_party/eipl && \
pip install -e .
## Diffusion Policy
RUN apt update -q -qq && \
apt install -q -qq -y libosmesa6-dev libglfw3 patchelf && \
apt clean && \
rm -rf /var/lib/apt/lists/
# libgl1-mesa-glx is not required in irsl_system
RUN source /irsl_venv/bin/activate && \
cd /RoboManipBaselines && \
pip install -e .[diffusion-policy] && \
cd third_party/diffusion_policy && \
pip install -e .
### patched by IRSL
RUN <<EOF
cd /RoboManipBaselines
cat - << _DOC_ | patch -p1
diff --git a/robo_manip_baselines/common/base/TrainBase.py b/robo_manip_baselines/common/base/TrainBase.py
index 5ef7ab7..4918b3d 100644
--- a/robo_manip_baselines/common/base/TrainBase.py
+++ b/robo_manip_baselines/common/base/TrainBase.py
@@ -168,7 +168,7 @@ class TrainBase(ABC):
)
parser.add_argument("--seed", type=int, default=42, help="random seed")
-
+ parser.add_argument("--save_interval", type=int, default=100, help="IRSL save_interval")
self.set_additional_args(parser)
if argv is None:
diff --git a/robo_manip_baselines/policy/act/TrainAct.py b/robo_manip_baselines/policy/act/TrainAct.py
index 78655d4..6140c10 100644
--- a/robo_manip_baselines/policy/act/TrainAct.py
+++ b/robo_manip_baselines/policy/act/TrainAct.py
@@ -97,8 +97,8 @@ class TrainAct(TrainBase):
self.update_best_ckpt(epoch_summary)
# Save current checkpoint
- if epoch % max(self.args.num_epochs // 10, 1) == 0:
- self.save_current_ckpt(f"epoch{epoch:0>3}")
+ if epoch % min(self.args.save_interval, max(self.args.num_epochs // 10, 1)) == 0:
+ self.save_current_ckpt(f"epoch{epoch:0>5}")
# Save last checkpoint
self.save_current_ckpt("last")
diff --git a/robo_manip_baselines/policy/diffusion_policy/TrainDiffusionPolicy.py b/robo_manip_baselines/policy/diffusion_policy/TrainDiffusionPolicy.py
index 938e8db..5d3bac2 100644
--- a/robo_manip_baselines/policy/diffusion_policy/TrainDiffusionPolicy.py
+++ b/robo_manip_baselines/policy/diffusion_policy/TrainDiffusionPolicy.py
@@ -336,8 +336,8 @@ class TrainDiffusionPolicy(TrainBase):
policy.train()
# Save current checkpoint
- if epoch % max(self.args.num_epochs // 10, 1) == 0:
- self.save_current_ckpt(f"epoch{epoch:0>4}", policy=policy)
+ if epoch % min(self.args.save_interval, max(self.args.num_epochs // 10, 1)) == 0:
+ self.save_current_ckpt(f"epoch{epoch:0>5}", policy=policy)
# Save last checkpoint
self.save_current_ckpt("last", policy=policy)
diff --git a/robo_manip_baselines/policy/sarnn/TrainSarnn.py b/robo_manip_baselines/policy/sarnn/TrainSarnn.py
index 6f15a45..fc5c9fc 100644
--- a/robo_manip_baselines/policy/sarnn/TrainSarnn.py
+++ b/robo_manip_baselines/policy/sarnn/TrainSarnn.py
@@ -243,8 +243,8 @@ class TrainSarnn(TrainBase):
self.update_best_ckpt(epoch_summary)
# Save current checkpoint
- if epoch % max(self.args.num_epochs // 10, 1) == 0:
- self.save_current_ckpt(f"epoch{epoch:0>4}")
+ if epoch % min(self.args.save_interval, max(self.args.num_epochs // 10, 1)) == 0:
+ self.save_current_ckpt(f"epoch{epoch:0>5}")
# Save last checkpoint
self.save_current_ckpt("last")
_DOC_
EOF