-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathrun_e3_full.py
More file actions
60 lines (53 loc) · 2.09 KB
/
Copy pathrun_e3_full.py
File metadata and controls
60 lines (53 loc) · 2.09 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
"""
E3 FULL: Smoothness regularisation ablation, 20 seeds, paired per-seed.
Rescue or downgrade the +7.6pp claim (review C3/M2).
Budgets 50 and 200 demos, 20 seeds (42-61), paired delta per seed.
"""
import numpy as np
import torch
import time
from env.arm2d import Arm2DEnv
from experts.demo_generator import generate_dataset
from models.promp import RBFBasis, ProMPPredictor, train_promp, evaluate_promp_closed_loop
SEED = 42
N_BFS = 15
N_EVAL = 100
EPOCHS = 500
N_SEEDS = 20
def main():
out = open("results_e3.txt", "w")
def log(*a):
msg = " ".join(str(x) for x in a)
print(msg, flush=True)
out.write(msg + "\n")
out.flush()
log("E3: Smoothness regularisation, 20 seeds, paired")
log(f"Date: {time.strftime('%Y-%m-%d %H:%M')}")
env = Arm2DEnv()
basis = RBFBasis(N_BFS)
for b in [50, 200]:
log(f"\n=== Budget {b} demos ===")
deltas = []
for seed in range(N_SEEDS):
np.random.seed(SEED + seed)
torch.manual_seed(SEED + seed)
data = generate_dataset(env, b, noise_levels=[0.03, 0.06, 0.09] * (b // 3 + 1))
# baseline CL
m0 = ProMPPredictor(10, N_BFS, 2, with_phase=True)
train_promp(m0, data, basis, epochs=EPOCHS, verbose=False, closed_loop=True)
sr0, _, _ = evaluate_promp_closed_loop(m0, env, basis)
# smoothness CL
m1 = ProMPPredictor(10, N_BFS, 2, with_phase=True)
train_promp(m1, data, basis, epochs=EPOCHS, verbose=False, closed_loop=True,
torque_penalty_weight=0.001, penalty_type="smoothness")
sr1, _, _ = evaluate_promp_closed_loop(m1, env, basis)
deltas.append(sr1 - sr0)
log(f" seed{seed:2d}: base {sr0:.3f} | smooth {sr1:.3f} | delta {sr1-sr0:+.3f}")
d = np.array(deltas)
mu = d.mean()
sd = d.std(ddof=1)
t_stat = mu / (sd / np.sqrt(len(d))) if sd > 0 else float('nan')
log(f" SUMMARY: mean delta {mu:+.3f} ± {sd:.3f}, paired t={t_stat:.2f}, n={len(d)}")
out.close()
if __name__ == "__main__":
main()