-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathrun_e8.py
More file actions
89 lines (78 loc) · 3.62 KB
/
Copy pathrun_e8.py
File metadata and controls
89 lines (78 loc) · 3.62 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
"""
E8: OOD mirror-symmetry test. Distinguish distribution shift from task difficulty.
Original: train RIGHT half-plane, test LEFT (hard, arm starts right).
Mirror: train LEFT half-plane, test RIGHT (easier, same side as arm).
If both directions collapse OOD -> distribution shift. If easy-side OOD much better -> difficulty confound.
Also evaluates CL (original run_generalization omitted it) and reports real OOD numbers.
"""
import numpy as np
import torch
import time
from env.arm2d import Arm2DEnv
from experts.demo_generator import generate_dataset
from models.bc import BCMLP, train_bc, evaluate_bc
from models.promp import RBFBasis, ProMPPredictor, train_promp, evaluate_promp_closed_loop
SEED = 42
N_BFS = 15
N_EVAL = 100
N_DEMOS = 50
N_SEEDS = 5
EPOCHS = 200
def sample_half(env, left):
r = np.random.uniform(0.1, env.l1 + env.l2 - 0.05)
if left:
theta = np.random.uniform(np.pi / 2, 3 * np.pi / 2)
else:
theta = np.random.uniform(-np.pi / 2, np.pi / 2)
return np.array([r * np.cos(theta), r * np.sin(theta)])
def main():
out = open("results_e8.txt", "w")
def log(*a):
msg = " ".join(str(x) for x in a)
print(msg, flush=True)
out.write(msg + "\n")
out.flush()
log("E8: OOD mirror-symmetry — distribution shift vs task difficulty")
log(f"Date: {time.strftime('%Y-%m-%d %H:%M')}")
env = Arm2DEnv()
basis = RBFBasis(N_BFS)
# Direction A: train RIGHT, test LEFT (original, hard OOD)
# Direction B: train LEFT, test RIGHT (mirror, easy OOD)
for direction, train_left, test_left, label in [
("A: train RIGHT -> test LEFT (hard)", False, True, "A"),
("B: train LEFT -> test RIGHT (easy)", True, False, "B"),
]:
log(f"\n=== {label}: {direction} ===")
log(f"{'Method':8s} | {'in-dist':>12} | {'OOD':>12} | {'gap':>8}")
for method in ["BC", "CL"]:
ind_s, ood_s = [], []
for seed in range(N_SEEDS):
np.random.seed(SEED + seed)
torch.manual_seed(SEED + seed)
# train data in train_half
orig = env._sample_reachable_target
env._sample_reachable_target = lambda: sample_half(env, train_left)
data = generate_dataset(env, N_DEMOS, noise_levels=[0.03, 0.06, 0.09] * (N_DEMOS // 3 + 1))
env._sample_reachable_target = orig
if not data:
continue
if method == "BC":
m = BCMLP(10, 2)
train_bc(m, data, epochs=EPOCHS, verbose=False)
sr_ind, _, _ = evaluate_bc(m, env, n_episodes=N_EVAL)
env._sample_reachable_target = lambda: sample_half(env, test_left)
sr_ood, _, _ = evaluate_bc(m, env, n_episodes=N_EVAL)
env._sample_reachable_target = orig
else:
m = ProMPPredictor(10, N_BFS, 2, with_phase=True)
train_promp(m, data, basis, epochs=EPOCHS, verbose=False, closed_loop=True)
sr_ind, _, _ = evaluate_promp_closed_loop(m, env, basis, n_episodes=N_EVAL)
env._sample_reachable_target = lambda: sample_half(env, test_left)
sr_ood, _, _ = evaluate_promp_closed_loop(m, env, basis, n_episodes=N_EVAL)
env._sample_reachable_target = orig
ind_s.append(sr_ind)
ood_s.append(sr_ood)
log(f"{method:8s} | {np.mean(ind_s):.3f}±{np.std(ind_s):.3f} | {np.mean(ood_s):.3f}±{np.std(ood_s):.3f} | {np.mean(ind_s)-np.mean(ood_s):+.3f}")
out.close()
if __name__ == "__main__":
main()