tests/test.sh exit code (0 → resolved); the classification below is post-hoc and cannot change it./app/objectives.py
1 from __future__ import annotations
2
3 import torch
4
5
6 def masked_mean(values, mask, axis=None):
7 raise NotImplementedError
8
9
10 def masked_sum(values, mask, axis=None):
11 raise NotImplementedError
12
13
14 def logsumexp(x, axis):
15 raise NotImplementedError
16
17
18 def log_softmax(x, axis):
19 raise NotImplementedError
20
21
22 def selective_logprobs(logits, labels, mask):
23 raise NotImplementedError
24
25
26 def token_logprobs(logits, labels):
27 raise NotImplementedError
28
29
30 def sequence_logprob(logits, labels, mask, length_normalize):
31 raise NotImplementedError
32
33
34 def entropy(logits, mask):
35 raise NotImplementedError
36
37
38 def dpo_loss(pc, pr, rc, rr, beta, label_smoothing):
39 raise NotImplementedError
40
41
42 def ipo_loss(pc, pr, rc, rr, beta):
43 raise NotImplementedError
44
45
46 def grpo_advantages(rewards, group_size, scale_by_std):
47 raise NotImplementedError
48
49
50 def gae(rewards, values, next_value, gamma, lam):
51 raise NotImplementedError
52
53
54 def kl_penalty(logp, ref_logp, estimator):
55 raise NotImplementedError
56
57
58 def clipped_pg_loss(logp, old_logp, advantages, mask, clip_low, clip_high):
59 raise NotImplementedError
60
61
62 def value_loss(values, old_values, returns, clip):
63 raise NotImplementedError
64
65
66 def whiten(values, mask, shift_mean):
67 raise NotImplementedError
68
69
70 def dpo_sequence_loss(pc_logits, pr_logits, rc_logits, rr_logits,
71 chosen_labels, rejected_labels, chosen_mask, rejected_mask,
72 beta, label_smoothing):
73 raise NotImplementedError
74
75
76 def grpo_objective(logits, old_logits, ref_logits, labels, completion_mask,
77 rewards, group_size, beta, clip_low, clip_high, scale_by_std,
78 kl_estimator):
79 raise NotImplementedError
80
81
82 def ppo_objective(rewards, values, old_values, logp, old_logp, next_value,
83 gamma, lam, clip_low, clip_high, vf_clip, vf_coef):
84 raise NotImplementedError
85
86
87 def rloo_advantages(rewards, group_size):
88 raise NotImplementedError
89
90
91 def reverse_kl(logp, ref_logp):
92 raise NotImplementedError
93
94
95 def importance_ratio(logp, old_logp, clip):
96 raise NotImplementedError
97
98
99 def discounted_returns(rewards, gamma):
100 raise NotImplementedError
101
102
103 def normalize(x, eps):
104 raise NotImplementedError
105
106
107 def top_p_mask(probs, p):
108 raise NotImplementedError
109
110
111 def smoothed_nll(logits, labels, smoothing):
112 raise NotImplementedError
113
114
115 def bradley_terry_logit(chosen_reward, rejected_reward, beta):
116 raise NotImplementedError
117
118
119 def rloo_objective(logits, old_logits, labels, mask, rewards, group_size, clip_low, clip_high):
120 raise NotImplementedError
121
122
123 def reverse_kl_objective(logits, ref_logits, labels, mask, advantages, beta):
124 raise NotImplementedError
125
126
127 def cross_entropy(logits, labels, ignore_index):
128 raise NotImplementedError
129
130
131 def top_k_mask(logits, k):
132 raise NotImplementedError
133
134
135 def group_mean_baseline(rewards, group_size):
136 raise NotImplementedError
137
138
139 def lambda_returns(rewards, values, next_value, gamma, lam):
140 raise NotImplementedError
141
142
143 def symmetric_kl(logp, ref_logp):
144 raise NotImplementedError
145
146
147 def huber_value_loss(values, returns, delta):
148 raise NotImplementedError
149
150
151 def normalized_entropy(logits, mask):
152 raise NotImplementedError
153
154
155 def clip_fraction(logp, old_logp, clip):
156 raise NotImplementedError
157
158
159 def masked_whiten(values, mask, shift_mean):
160 raise NotImplementedError
161
162
163 def logprob_at_temperature(logits, labels, temperature):
164 raise NotImplementedError
165
166
167 def advantage_mean_std(advantages, mask):
168 raise NotImplementedError
169
170
171 def argmax_tokens(logits):
172 raise NotImplementedError
173
174
175 def mode_label(labels):
176 raise NotImplementedError
177
178
179 def median_reward(rewards):
180 raise NotImplementedError
181
182
183 def masked_max(values, mask):
184 raise NotImplementedError
185
186
187 def quantile_lower(x, q):
188 raise NotImplementedError
189
190
191 def pad_mask_from_lengths(lengths, max_len):
192 raise NotImplementedError
193
194
195 def first_nonzero_index(mask):
196 raise NotImplementedError
197
198
199 def cumulative_max(x):
200 raise NotImplementedError
201
202
203 def bucketize_reward(rewards, boundaries):
204 raise NotImplementedError
205
1 from __future__ import annotations
2
3 import torch
4
5 import reforacle
6
7 g = torch.Generator().manual_seed(0)
8 values = torch.randn(8, generator=g)
9 mask = (torch.rand(8, generator=g) > 0.5).float()
10
11 print("inputs:")
12 print(" values =", values.tolist())
13 print(" mask =", mask.tolist())
14 print("reference masked_mean ->", reforacle.masked_mean(values, mask).item())
15 print("reference masked_mean (empty mask) ->", reforacle.masked_mean(torch.randn(4, generator=g), torch.zeros(4)).item())
16 print()
17 print("Probe any reference function this way to discover its exact numerical")
18 print("convention, then implement it in /app/objectives.py to match.")
19
Remaining time (hours:minutes):
1:59
[stdout]
Remaining time (hours:minutes):
1:59/app/probe.py
1 from __future__ import annotations
2
3 import torch
4
5
6 def masked_mean(values, mask, axis=None):
7 raise NotImplementedError
8
9
10 def masked_sum(values, mask, axis=None):
11 raise NotImplementedError
12
13
14 def logsumexp(x, axis):
15 raise NotImplementedError
16
17
18 def log_softmax(x, axis):
19 raise NotImplementedError
20
21
22 def selective_logprobs(logits, labels, mask):
23 raise NotImplementedError
24
25
26 def token_logprobs(logits, labels):
27 raise NotImplementedError
28
29
30 def sequence_logprob(logits, labels, mask, length_normalize):
31 raise NotImplementedError
32
33
34 def entropy(logits, mask):
35 raise NotImplementedError
36
37
38 def dpo_loss(pc, pr, rc, rr, beta, label_smoothing):
39 raise NotImplementedError
40
41
42 def ipo_loss(pc, pr, rc, rr, beta):
43 raise NotImplementedError
44
45
46 def grpo_advantages(rewards, group_size, scale_by_std):
47 raise NotImplementedError
48
49
50 def gae(rewards, values, next_value, gamma, lam):
51 raise NotImplementedError
52
53
54 def kl_penalty(logp, ref_logp, estimator):
55 raise NotImplementedError
56
57
58 def clipped_pg_loss(logp, old_logp, advantages, mask, clip_low, clip_high):
59 raise NotImplementedError
60
61
62 def value_loss(values, old_values, returns, clip):
63 raise NotImplementedError
64
65
66 def whiten(values, mask, shift_mean):
67 raise NotImplementedError
68
69
70 def dpo_sequence_loss(pc_logits, pr_logits, rc_logits, rr_logits,
71 chosen_labels, rejected_labels, chosen_mask, rejected_mask,
72 beta, label_smoothing):
73 raise NotImplementedError
74
75
76 def grpo_objective(logits, old_logits, ref_logits, labels, completion_mask,
77 rewards, group_size, beta, clip_low, clip_high, scale_by_std,
78 kl_estimator):
79 raise NotImplementedError
80
81
82 def ppo_objective(rewards, values, old_values, logp, old_logp, next_value,
83 gamma, lam, clip_low, clip_high, vf_clip, vf_coef):
84 raise NotImplementedError
85
86
87 def rloo_advantages(rewards, group_size):
88 raise NotImplementedError
89
90
91 def reverse_kl(logp, ref_logp):
92 raise NotImplementedError
93
94
95 def importance_ratio(logp, old_logp, clip):
96 raise NotImplementedError
97
98
99 def discounted_returns(rewards, gamma):
100 raise NotImplementedError
101
102
103 def normalize(x, eps):
104 raise NotImplementedError
105
106
107 def top_p_mask(probs, p):
108 raise NotImplementedError
109
110
111 def smoothed_nll(logits, labels, smoothing):
112 raise NotImplementedError
113
114
115 def bradley_terry_logit(chosen_reward, rejected_reward, beta):
116 raise NotImplementedError
117
118
119 def rloo_objective(logits, old_logits, labels, mask, rewards, group_size, clip_low, clip_high):
120 raise NotImplementedError
121
122
123 def reverse_kl_objective(logits, ref_logits, labels, mask, advantages, beta):
124 raise NotImplementedError
125
126
127 def cross_entropy(logits, labels, ignore_index):
128 raise NotImplementedError
129
130
131 def top_k_mask(logits, k):
132 raise NotImplementedError
133
134
135 def group_mean_baseline(rewards, group_size):
136 raise NotImplementedError
137
138
139 def lambda_returns(rewards, values, next_value, gamma, lam):
140 raise NotImplementedError
141
142
143 def symmetric_kl(logp, ref_logp):
144 raise NotImplementedError
145
146
147 def huber_value_loss(values, returns, delta):
148 raise NotImplementedError
149
150
151 def normalized_entropy(logits, mask):
152 raise NotImplementedError
153
154
155 def clip_fraction(logp, old_logp, clip):
156 raise NotImplementedError
157
158
159 def masked_whiten(values, mask, shift_mean):
160 raise NotImplementedError
161
162
163 def logprob_at_temperature(logits, labels, temperature):
164 raise NotImplementedError
165
166
167 def advantage_mean_std(advantages, mask):
168 raise NotImplementedError
169
170
171 def argmax_tokens(logits):
172 raise NotImplementedError
173
174
175 def mode_label(labels):
176 raise NotImplementedError
177
178
179 def median_reward(rewards):
180 raise NotImplementedError
181
182
183 def masked_max(values, mask):
184 raise NotImplementedError
185
186
187 def quantile_lower(x, q):
188 raise NotImplementedError
189
190
191 def pad_mask_from_lengths(lengths, max_len):
192 raise NotImplementedError
193
194
195 def first_nonzero_index(mask):
196 raise NotImplementedError
197
198
199 def cumulative_max(x):
200 raise NotImplementedError
201
202
203 def bucketize_reward(rewards, boundaries):
204 raise NotImplementedError
205
1 from __future__ import annotations
2
3 import torch
4
5 import reforacle
6
7 g = torch.Generator().manual_seed(0)
8 values = torch.randn(8, generator=g)
9 mask = (torch.rand(8, generator=g) > 0.5).float()
10
11 print("inputs:")
12 print(" values =", values.tolist())
13 print(" mask =", mask.tolist())
14 print("reference masked_mean ->", reforacle.masked_mean(values, mask).item())
15 print("reference masked_mean (empty mask) ->", reforacle.masked_mean(torch.randn(4, generator=g), torch.zeros(4)).item())
16 print()
17 print("Probe any reference function this way to discover its exact numerical")
18 print("convention, then implement it in /app/objectives.py to match.")
19
Remaining time (hours:minutes):
1:59
[stdout]
Remaining time (hours:minutes):
1:59Check remaining budget
bash /app/timer.sh
1 from __future__ import annotations
2
3 import torch
4
5
6 def masked_mean(values, mask, axis=None):
7 raise NotImplementedError
8
9
10 def masked_sum(values, mask, axis=None):
11 raise NotImplementedError
12
13
14 def logsumexp(x, axis):
15 raise NotImplementedError
16
17
18 def log_softmax(x, axis):
19 raise NotImplementedError
20
21
22 def selective_logprobs(logits, labels, mask):
23 raise NotImplementedError
24
25
26 def token_logprobs(logits, labels):
27 raise NotImplementedError
28
29
30 def sequence_logprob(logits, labels, mask, length_normalize):
31 raise NotImplementedError
32
33
34 def entropy(logits, mask):
35 raise NotImplementedError
36
37
38 def dpo_loss(pc, pr, rc, rr, beta, label_smoothing):
39 raise NotImplementedError
40
41
42 def ipo_loss(pc, pr, rc, rr, beta):
43 raise NotImplementedError
44
45
46 def grpo_advantages(rewards, group_size, scale_by_std):
47 raise NotImplementedError
48
49
50 def gae(rewards, values, next_value, gamma, lam):
51 raise NotImplementedError
52
53
54 def kl_penalty(logp, ref_logp, estimator):
55 raise NotImplementedError
56
57
58 def clipped_pg_loss(logp, old_logp, advantages, mask, clip_low, clip_high):
59 raise NotImplementedError
60
61
62 def value_loss(values, old_values, returns, clip):
63 raise NotImplementedError
64
65
66 def whiten(values, mask, shift_mean):
67 raise NotImplementedError
68
69
70 def dpo_sequence_loss(pc_logits, pr_logits, rc_logits, rr_logits,
71 chosen_labels, rejected_labels, chosen_mask, rejected_mask,
72 beta, label_smoothing):
73 raise NotImplementedError
74
75
76 def grpo_objective(logits, old_logits, ref_logits, labels, completion_mask,
77 rewards, group_size, beta, clip_low, clip_high, scale_by_std,
78 kl_estimator):
79 raise NotImplementedError
80
81
82 def ppo_objective(rewards, values, old_values, logp, old_logp, next_value,
83 gamma, lam, clip_low, clip_high, vf_clip, vf_coef):
84 raise NotImplementedError
85
86
87 def rloo_advantages(rewards, group_size):
88 raise NotImplementedError
89
90
91 def reverse_kl(logp, ref_logp):
92 raise NotImplementedError
93
94
95 def importance_ratio(logp, old_logp, clip):
96 raise NotImplementedError
97
98
99 def discounted_returns(rewards, gamma):
100 raise NotImplementedError
101
102
103 def normalize(x, eps):
104 raise NotImplementedError
105
106
107 def top_p_mask(probs, p):
108 raise NotImplementedError
109
110
111 def smoothed_nll(logits, labels, smoothing):
112 raise NotImplementedError
113
114
115 def bradley_terry_logit(chosen_reward, rejected_reward, beta):
116 raise NotImplementedError
117
118
119 def rloo_objective(logits, old_logits, labels, mask, rewards, group_size, clip_low, clip_high):
120 raise NotImplementedError
121
122
123 def reverse_kl_objective(logits, ref_logits, labels, mask, advantages, beta):
124 raise NotImplementedError
125
126
127 def cross_entropy(logits, labels, ignore_index):
128 raise NotImplementedError
129
130
131 def top_k_mask(logits, k):
132 raise NotImplementedError
133
134
135 def group_mean_baseline(rewards, group_size):
136 raise NotImplementedError
137
138
139 def lambda_returns(rewards, values, next_value, gamma, lam):
140 raise NotImplementedError
141
142
143 def symmetric_kl(logp, ref_logp):
144 raise NotImplementedError
145
146
147 def huber_value_loss(values, returns, delta):
148 raise NotImplementedError
149
150
151 def normalized_entropy(logits, mask):
152 raise NotImplementedError
153
154
155 def clip_fraction(logp, old_logp, clip):
156 raise NotImplementedError
157
158
159 def masked_whiten(values, mask, shift_mean):
160 raise NotImplementedError
161
162
163 def logprob_at_temperature(logits, labels, temperature):
164 raise NotImplementedError
165
166
167 def advantage_mean_std(advantages, mask):
168 raise NotImplementedError
169
170
171 def argmax_tokens(logits):
172 raise NotImplementedError
173
174
175 def mode_label(labels):
176 raise NotImplementedError
177
178
179 def median_reward(rewards):
180 raise NotImplementedError
181
182
183 def masked_max(values, mask):
184 raise NotImplementedError
185
186
187 def quantile_lower(x, q):
188 raise NotImplementedError
189
190
191 def pad_mask_from_lengths(lengths, max_len):
192 raise NotImplementedError
193
194
195 def first_nonzero_index(mask):
196 raise NotImplementedError
197
198
199 def cumulative_max(x):
200 raise NotImplementedError
201
202
203 def bucketize_reward(rewards, boundaries):
204 raise NotImplementedError
205
1 from __future__ import annotations
2
3 import torch
4
5 import reforacle
6
7 g = torch.Generator().manual_seed(0)
8 values = torch.randn(8, generator=g)
9 mask = (torch.rand(8, generator=g) > 0.5).float()
10
11 print("inputs:")
12 print(" values =", values.tolist())
13 print(" mask =", mask.tolist())
14 print("reference masked_mean ->", reforacle.masked_mean(values, mask).item())
15 print("reference masked_mean (empty mask) ->", reforacle.masked_mean(torch.randn(4, generator=g), torch.zeros(4)).item())
16 print()
17 print("Probe any reference function this way to discover its exact numerical")
18 print("convention, then implement it in /app/objectives.py to match.")
19
Remaining time (hours:minutes):
1:59
[stdout]
Remaining time (hours:minutes):
1:59Run probe to confirm reforacle works
python3 probe.py
inputs: values = [1.5409960746765137, -0.293428897857666, -2.1787893772125244, 0.5684312582015991, -1.0845223665237427, -1.3985954523086548, 0.40334683656692505, 0.8380263447761536] mask = [1.0, 1.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0] reference masked_mean -> -0.07950486242771149 reference masked_mean (empty mask) -> 0.0 Probe any reference function this way to discover its exact numerical convention, then implement it in /app/objectives.py to match. [stdout] inputs: values = [1.5409960746765137, -0.293428897857666, -2.1787893772125244, 0.5684312582015991, -1.0845223665237427, -1.3985954523086548, 0.40334683656692505, 0.8380263447761536] mask = [1.0, 1.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0] reference masked_mean -> -0.07950486242771149 reference masked_mean (empty mask) -> 0.0 Probe any reference function this way to discover its exact numerical convention, then implement it in /app/objectives.py to match.
Probe primitives
cat > /tmp/p1.py << 'EOF'
import torch, reforacle as R
g=torch.Generator().manual_seed(1)
def rn(*s): return torch.randn(*s, generator=g)
# masked_mean axis
v=rn(3,4); m=(torch.rand(3,4,generator=g)>0.4).float()
print("mm none", R.masked_mean(v,m))
print("mm ax0", R.masked_mean(v,m,axis=0))
print("mm ax1", R.masked_mean(v,m,axis=1))
print("ms none", R.masked_sum(v,m))
print("ms ax1", R.masked_sum(v,m,axis=1))
# logsumexp/log_softmax
x=rn(2,5)
print("lse", R.logsumexp(x,axis=1))
print("lsm", R.log_softmax(x,axis=1))
# token_logprobs
logits=rn(2,3,5); labels=torch.randint(0,5,(2,3),generator=g)
print("tlp", R.token_logprobs(logits,labels))
mask=(torch.rand(2,3,generator=g)>0.3).float()
print("slp", R.selective_logprobs(logits,labels,mask))
print("seqlp F", R.sequence_logprob(logits,labels,mask,False))
print("seqlp T", R.sequence_logprob(logits,labels,mask,True))
print("entropy", R.entropy(logits,mask))
print("normalize", R.normalize(rn(6),1e-8))
EOF
python3 /tmp/p1.pyExit code 1
Traceback (most recent call last):
File "/tmp/p1.py", line 1, in <module>
import torch, reforacle as R
ModuleNotFoundError: No module named 'reforacle'
[error] tool reported failureRun probe from /app
cp /tmp/p1.py /app/_p1.py && python3 _p1.py
mm none tensor(-0.1978)
mm ax0 tensor([-1.0276, 0.0504, -0.7305, 0.5015])
mm ax1 tensor([ 0.3166, -0.4357, -1.0276])
ms none tensor(-1.3849)
ms ax1 tensor([ 0.9499, -1.3072, -1.0276])
lse tensor([3.7453, 2.9333])
lsm tensor([[-5.3180, -3.8686, -0.1584, -5.5766, -2.1466],
[-4.2103, -2.6078, -3.4124, -1.5543, -0.4047]])
tlp tensor([[-2.6539, -1.9474, -2.8140],
[-1.6476, -1.2624, -3.1213]])
slp tensor([-7.4154, -2.9100])
seqlp F tensor([-7.4154, -2.9100])
seqlp T tensor([-2.4718, -1.4550])
entropy tensor(1.3433)
normalize tensor([-1.8910, -0.0302, 0.8794, 0.6768, 0.4327, -0.0676])
[stdout]
mm none tensor(-0.1978)
mm ax0 tensor([-1.0276, 0.0504, -0.7305, 0.5015])
mm ax1 tensor([ 0.3166, -0.4357, -1.0276])
ms none tensor(-1.3849)
ms ax1 tensor([ 0.9499, -1.3072, -1.0276])
lse tensor([3.7453, 2.9333])
lsm tensor([[-5.3180, -3.8686, -0.1584, -5.5766, -2.1466],
[-4.2103, -2.6078, -3.4124, -1.5543, -0.4047]])
tlp tensor([[-2.6539, -1.9474, -2.8140],
[-1.6476, -1.2624, -3.1213]])
slp tensor([-7.4154, -2.9100])
seqlp F tensor([-7.4154, -2.9100])
seqlp T tensor([-2.4718, -1.4550])
entropy tensor(1.3433)
normalize tensor([-1.8910, -0.0302, 0.8794, 0.6768, 0.4327, -0.0676])Probe normalize/entropy conventions
cat > /app/_p2.py << 'EOF'
import torch, reforacle as R
x=torch.tensor([1.,2.,3.,4.])
print("normalize eps1e-8", R.normalize(x,1e-8))
mu=x.mean();
print(" (x-mu)/std(unbiased)", (x-mu)/x.std(unbiased=True))
print(" (x-mu)/std(biased)", (x-mu)/x.std(unbiased=False))
print(" (x-mu)/(std_b+eps)", (x-mu)/(x.std(unbiased=False)+1e-8))
# entropy check with simple logits
logits=torch.tensor([[[0.,0.],[10.,-10.]]]) # 1,2,2
mask=torch.tensor([[1.,1.]])
p=torch.softmax(logits,-1); ent=-(p*torch.log_softmax(logits,-1)).sum(-1)
print("per token ent", ent)
print("R.entropy full mask", R.entropy(logits,mask))
mask2=torch.tensor([[1.,0.]])
print("R.entropy half mask", R.entropy(logits,mask2))
# masked_mean empty axis
v=torch.tensor([[1.,2.],[3.,4.]]); m=torch.tensor([[0.,0.],[1.,1.]])
print("mm ax1 empty row", R.masked_mean(v,m,axis=1))
EOF
python3 /app/_p2.pynormalize eps1e-8 tensor([-1.1619, -0.3873, 0.3873, 1.1619]) (x-mu)/std(unbiased) tensor([-1.1619, -0.3873, 0.3873, 1.1619]) (x-mu)/std(biased) tensor([-1.3416, -0.4472, 0.4472, 1.3416]) (x-mu)/(std_b+eps) tensor([-1.3416, -0.4472, 0.4472, 1.3416]) per token ent tensor([[6.9315e-01, 4.1223e-08]]) R.entropy full mask tensor(0.3466) R.entropy half mask tensor(0.6931) mm ax1 empty row tensor([0.0000, 3.5000]) [stdout] normalize eps1e-8 tensor([-1.1619, -0.3873, 0.3873, 1.1619]) (x-mu)/std(unbiased) tensor([-1.1619, -0.3873, 0.3873, 1.1619]) (x-mu)/std(biased) tensor([-1.3416, -0.4472, 0.4472, 1.3416]) (x-mu)/(std_b+eps) tensor([-1.3416, -0.4472, 0.4472, 1.3416]) per token ent tensor([[6.9315e-01, 4.1223e-08]]) R.entropy full mask tensor(0.3466) R.entropy half mask tensor(0.6931) mm ax1 empty row tensor([0.0000, 3.5000])
Probe KL and RL functions
cat > /app/_p3.py << 'EOF'
import torch, reforacle as R
lp=torch.tensor([-1.0,-2.0,-0.5]); rlp=torch.tensor([-1.5,-1.0,-0.7])
print("reverse_kl", R.reverse_kl(lp,rlp))
print("symmetric_kl", R.symmetric_kl(lp,rlp))
for est in ["k1","k2","k3"]:
try: print("kl_penalty",est, R.kl_penalty(lp,rlp,est))
except Exception as e: print(est,"ERR",e)
print("importance_ratio noclip", R.importance_ratio(lp,rlp,None))
print("importance_ratio clip0.1", R.importance_ratio(lp,rlp,0.1))
r=torch.tensor([1.0,2.0,3.0])
print("discounted_returns g0.9", R.discounted_returns(r,0.9))
# gae
rewards=torch.tensor([1.0,0.0,-1.0]); values=torch.tensor([0.5,0.2,0.1])
print("gae", R.gae(rewards,values,torch.tensor(0.3),0.99,0.95))
print("lambda_returns", R.lambda_returns(rewards,values,torch.tensor(0.3),0.99,0.95))
v=torch.tensor([1.,2.,3.,4.,5.]); m=torch.tensor([1.,1.,0.,1.,1.])
print("whiten shiftT", R.whiten(v,m,True))
print("whiten shiftF", R.whiten(v,m,False))
print("masked_whiten shiftT", R.masked_whiten(v,m,True))
print("masked_whiten shiftF", R.masked_whiten(v,m,False))
EOF
python3 /app/_p3.pyreverse_kl tensor([0.1065, 0.7183, 0.0187]) symmetric_kl tensor([0.1276, 0.5431, 0.0201]) kl_penalty k1 tensor([ 0.5000, -1.0000, 0.2000]) kl_penalty k2 tensor([0.1250, 0.5000, 0.0200]) kl_penalty k3 tensor([0.1065, 0.7183, 0.0187]) importance_ratio noclip tensor([1.6487, 0.3679, 1.2214]) importance_ratio clip0.1 tensor([1.1000, 0.9000, 1.1000]) discounted_returns g0.9 tensor([5.2300, 4.7000, 3.0000]) gae tensor([-0.1073, -0.8562, -0.8030]) lambda_returns tensor([ 0.3927, -0.6562, -0.7030]) whiten shiftT tensor([-1.2649, -0.6325, 0.0000, 0.6325, 1.2649]) whiten shiftF tensor([1.7351, 2.3675, 3.0000, 3.6325, 4.2649]) masked_whiten shiftT tensor([-1.2649, -0.6325, 0.0000, 0.6325, 1.2649]) masked_whiten shiftF tensor([1.7351, 2.3675, 3.0000, 3.6325, 4.2649]) [stdout] reverse_kl tensor([0.1065, 0.7183, 0.0187]) symmetric_kl tensor([0.1276, 0.5431, 0.0201]) kl_penalty k1 tensor([ 0.5000, -1.0000, 0.2000]) kl_penalty k2 tensor([0.1250, 0.5000, 0.0200]) kl_penalty k3 tensor([0.1065, 0.7183, 0.0187]) importance_ratio noclip tensor([1.6487, 0.3679, 1.2214]) importance_ratio clip0.1 tensor([1.1000, 0.9000, 1.1000]) discounted_returns g0.9 tensor([5.2300, 4.7000, 3.0000]) gae tensor([-0.1073, -0.8562, -0.8030]) lambda_returns tensor([ 0.3927, -0.6562, -0.7030]) whiten shiftT tensor([-1.2649, -0.6325, 0.0000, 0.6325, 1.2649]) whiten shiftF tensor([1.7351, 2.3675, 3.0000, 3.6325, 4.2649]) masked_whiten shiftT tensor([-1.2649, -0.6325, 0.0000, 0.6325, 1.2649]) masked_whiten shiftF tensor([1.7351, 2.3675, 3.0000, 3.6325, 4.2649])
Probe whiten/pg/value losses
cat > /app/_p4.py << 'EOF'
import torch, reforacle as R
v=torch.tensor([1.,2.,3.,10.,20.]); m=torch.tensor([1.,1.,1.,0.,0.])
print("whiten T", R.whiten(v,m,True))
print("masked_whiten T", R.masked_whiten(v,m,True))
# manual masked
mv=v[m.bool()];
print(" masked mean/std", mv.mean(), mv.std(unbiased=True), mv.std(unbiased=False))
print(" full mean/std", v.mean(), v.std(unbiased=True))
# clipped_pg_loss
logp=torch.tensor([-1.0,-2.0,-0.5,-1.2]); old=torch.tensor([-1.1,-1.5,-0.5,-1.0])
adv=torch.tensor([1.0,-1.0,2.0,-0.5]); mask=torch.tensor([1.,1.,1.,0.])
print("clipped_pg", R.clipped_pg_loss(logp,old,adv,mask,0.2,0.2))
print("clip_fraction", R.clip_fraction(logp,old,0.2))
# value_loss
values=torch.tensor([1.,2.,3.]); oldv=torch.tensor([1.1,1.8,3.2]); ret=torch.tensor([1.5,2.5,2.0])
print("value_loss clip0.2", R.value_loss(values,oldv,ret,0.2))
print("value_loss clipNone", R.value_loss(values,oldv,ret,None))
print("huber", R.huber_value_loss(values,ret,1.0))
EOF
python3 /app/_p4.pyExit code 1
whiten T tensor([-1.2247, 0.0000, 1.2247, 9.7980, 22.0454])
masked_whiten T tensor([-1.2247, 0.0000, 1.2247, 0.0000, 0.0000])
masked mean/std tensor(2.) tensor(1.) tensor(0.8165)
full mean/std tensor(7.2000) tensor(7.9812)
clipped_pg tensor(-0.7684)
clip_fraction tensor(0.2500)
value_loss clip0.2 tensor(0.2500)
Traceback (most recent call last):
File "/app/_p4.py", line 17, in <module>
print("value_loss clipNone", R.value_loss(values,oldv,ret,None))
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "reforacle.py", line 109, in reforacle.value_loss
TypeError: bad operand type for unary -: 'NoneType'
[error] tool reported failureProbe huber and group advantages
cat > /app/_p5.py << 'EOF'
import torch, reforacle as R
values=torch.tensor([1.,2.,3.]); ret=torch.tensor([1.5,2.5,5.0])
print("huber d1", R.huber_value_loss(values,ret,1.0))
# manual huber: 0.5*mean of huber
d=values-ret; a=d.abs(); delta=1.0
h=torch.where(a<=delta,0.5*d*d,delta*(a-0.5*delta))
print(" mean huber", h.mean(), "0.5*mean", 0.5*h.mean())
# group advantages
rewards=torch.tensor([1.,2.,3.,4.,5.,6.])
print("grpo scaleT", R.grpo_advantages(rewards,3,True))
print("grpo scaleF", R.grpo_advantages(rewards,3,False))
print("rloo", R.rloo_advantages(rewards,3))
print("group_mean_baseline", R.group_mean_baseline(rewards,3))
adv=torch.tensor([1.,2.,3.,4.]); m=torch.tensor([1.,1.,0.,1.])
print("advantage_mean_std", R.advantage_mean_std(adv,m))
EOF
python3 /app/_p5.pyhuber d1 tensor(0.5833) mean huber tensor(0.5833) 0.5*mean tensor(0.2917) grpo scaleT tensor([-1.2247, 0.0000, 1.2247, -1.2247, 0.0000, 1.2247]) grpo scaleF tensor([-1., 0., 1., -1., 0., 1.]) rloo tensor([-1.5000, 0.0000, 1.5000, -1.5000, 0.0000, 1.5000]) group_mean_baseline tensor([-1., 0., 1., -1., 0., 1.]) advantage_mean_std tensor([2.3333, 1.2472]) [stdout] huber d1 tensor(0.5833) mean huber tensor(0.5833) 0.5*mean tensor(0.2917) grpo scaleT tensor([-1.2247, 0.0000, 1.2247, -1.2247, 0.0000, 1.2247]) grpo scaleF tensor([-1., 0., 1., -1., 0., 1.]) rloo tensor([-1.5000, 0.0000, 1.5000, -1.5000, 0.0000, 1.5000]) group_mean_baseline tensor([-1., 0., 1., -1., 0., 1.]) advantage_mean_std tensor([2.3333, 1.2472])
Probe dpo/ce/entropy
cat > /app/_p6.py << 'EOF'
import torch, reforacle as R
r=R.advantage_mean_std(torch.tensor([1.,2.,4.]),torch.tensor([1.,1.,1.]))
print("ams type", type(r), r)
# dpo/ipo
pc,pr,rc,rr=torch.tensor([-2.]),torch.tensor([-3.]),torch.tensor([-2.5]),torch.tensor([-2.2])
print("dpo ls0", R.dpo_loss(pc,pr,rc,rr,0.1,0.0))
print("dpo ls0.1", R.dpo_loss(pc,pr,rc,rr,0.1,0.1))
print("ipo", R.ipo_loss(pc,pr,rc,rr,0.1))
print("bt", R.bradley_terry_logit(torch.tensor([1.0]),torch.tensor([0.3]),0.5))
# smoothed_nll / cross_entropy
logits=torch.tensor([[2.,1.,0.,-1.],[0.,0.,1.,2.]]); labels=torch.tensor([0,3])
print("smoothed0", R.smoothed_nll(logits,labels,0.0))
print("smoothed0.1", R.smoothed_nll(logits,labels,0.1))
print("ce", R.cross_entropy(logits,labels,-100))
labels2=torch.tensor([0,-100])
print("ce ignore", R.cross_entropy(logits,labels2,-100))
# normalized_entropy
lg=torch.tensor([[[0.,0.,0.,0.]]]); mk=torch.tensor([[1.]])
print("norm_ent uniform", R.normalized_entropy(lg,mk))
print("logprob_at_temp", R.logprob_at_temperature(logits.unsqueeze(0),labels.unsqueeze(0),2.0))
EOF
python3 /app/_p6.pyams type <class 'torch.Tensor'> tensor([2.3333, 1.2472]) dpo ls0 tensor(0.6303) dpo ls0.1 tensor(0.6433) ipo tensor(13.6900) bt tensor([0.3500]) smoothed0 tensor(0.4670) smoothed0.1 tensor(0.6045) ce tensor(0.4670) ce ignore tensor(0.4402) norm_ent uniform tensor(1.) logprob_at_temp tensor([[-0.7873, -0.8511]]) [stdout] ams type <class 'torch.Tensor'> tensor([2.3333, 1.2472]) dpo ls0 tensor(0.6303) dpo ls0.1 tensor(0.6433) ipo tensor(13.6900) bt tensor([0.3500]) smoothed0 tensor(0.4670) smoothed0.1 tensor(0.6045) ce tensor(0.4670) ce ignore tensor(0.4402) norm_ent uniform tensor(1.) logprob_at_temp tensor([[-0.7873, -0.8511]])
Probe masks and simple ops
cat > /app/_p7.py << 'EOF'
import torch, reforacle as R
probs=torch.tensor([[0.1,0.2,0.3,0.4],[0.5,0.3,0.15,0.05]])
print("top_p 0.7\n", R.top_p_mask(probs,0.7))
logits=torch.tensor([[2.,1.,0.,-1.],[0.,3.,1.,2.]])
print("top_k 2\n", R.top_k_mask(logits,2))
print("argmax", R.argmax_tokens(logits))
print("mode_label", R.mode_label(torch.tensor([1,2,2,3,3,3])))
print("median odd", R.median_reward(torch.tensor([3.,1.,2.,5.,4.])))
print("median even", R.median_reward(torch.tensor([1.,2.,3.,4.])))
v=torch.tensor([1.,5.,2.,8.]); m=torch.tensor([1.,0.,1.,1.])
print("masked_max", R.masked_max(v,m))
print("quantile_lower 0.25", R.quantile_lower(torch.tensor([1.,2.,3.,4.,5.]),0.25))
print("pad_mask", R.pad_mask_from_lengths(torch.tensor([2,3,1]),4))
print("first_nonzero", R.first_nonzero_index(torch.tensor([[0.,0.,1.,0.],[1.,0.,0.,0.],[0.,0.,0.,0.]])))
print("cummax", R.cumulative_max(torch.tensor([1.,3.,2.,5.,4.])))
print("bucketize", R.bucketize_reward(torch.tensor([0.5,1.5,2.5,3.5]),torch.tensor([1.,2.,3.])))
EOF
python3 /app/_p7.pytop_p 0.7
tensor([[0.0000, 0.0000, 0.4286, 0.5714],
[0.6250, 0.3750, 0.0000, 0.0000]])
top_k 2
tensor([[ True, True, False, False],
[False, True, False, True]])
argmax tensor([0, 1])
mode_label tensor(3)
median odd tensor(3.)
median even tensor(2.5000)
masked_max tensor(8.)
quantile_lower 0.25 tensor(2.)
pad_mask tensor([[1., 1., 0., 0.],
[1., 1., 1., 0.],
[1., 0., 0., 0.]])
first_nonzero tensor([ 2, 0, -1])
cummax tensor([1., 1., 3., 3., 5.])
bucketize tensor([0, 1, 2, 3])
[stdout]
top_p 0.7
tensor([[0.0000, 0.0000, 0.4286, 0.5714],
[0.6250, 0.3750, 0.0000, 0.0000]])
top_k 2
tensor([[ True, True, False, False],
[False, True, False, True]])
argmax tensor([0, 1])
mode_label tensor(3)
median odd tensor(3.)
median even tensor(2.5000)
masked_max tensor(8.)
quantile_lower 0.25 tensor(2.)
pad_mask tensor([[1., 1., 0., 0.],
[1., 1., 1., 0.],
[1., 0., 0., 0.]])
first_nonzero tensor([ 2, 0, -1])
cummax tensor([1., 1., 3., 3., 5.])
bucketize tensor([0, 1, 2, 3])Probe edge cases
cat > /app/_p8.py << 'EOF'
import torch, reforacle as R
print("cummax A", R.cumulative_max(torch.tensor([1.,3.,2.,5.,4.])))
print("cummax B", R.cumulative_max(torch.tensor([5.,1.,2.,3.])))
print("cummax 2d", R.cumulative_max(torch.tensor([[1.,3.,2.],[4.,0.,9.]])))
print("quantile 0.3 on 1..5", R.quantile_lower(torch.tensor([1.,2.,3.,4.,5.]),0.3))
print("quantile 0.5 on 1..4", R.quantile_lower(torch.tensor([1.,2.,3.,4.]),0.5))
print("bucketize eq", R.bucketize_reward(torch.tensor([1.,2.,3.]),torch.tensor([1.,2.,3.])))
print("mode tie", R.mode_label(torch.tensor([1,1,2,2])))
# top_p single dominant
print("top_p dom", R.top_p_mask(torch.tensor([[0.8,0.1,0.1]]),0.5))
print("top_p p1", R.top_p_mask(torch.tensor([[0.4,0.3,0.2,0.1]]),1.0))
EOF
python3 /app/_p8.pycummax A tensor([1., 1., 3., 3., 5.])
cummax B tensor([5., 5., 5., 5.])
cummax 2d tensor([[1., 3., 2.],
[1., 3., 2.]])
quantile 0.3 on 1..5 tensor(2.)
quantile 0.5 on 1..4 tensor(2.)
bucketize eq tensor([1, 2, 3])
mode tie tensor(2)
top_p dom tensor([[1., 0., 0.]])
top_p p1 tensor([[0.4000, 0.3000, 0.2000, 0.1000]])
[stdout]
cummax A tensor([1., 1., 3., 3., 5.])
cummax B tensor([5., 5., 5., 5.])
cummax 2d tensor([[1., 3., 2.],
[1., 3., 2.]])
quantile 0.3 on 1..5 tensor(2.)
quantile 0.5 on 1..4 tensor(2.)
bucketize eq tensor([1, 2, 3])
mode tie tensor(2)
top_p dom tensor([[1., 0., 0.]])
top_p p1 tensor([[0.4000, 0.3000, 0.2000, 0.1000]])Probe normalized_entropy, reverse_kl_objective
cat > /app/_p9.py << 'EOF'
import torch, reforacle as R
g=torch.Generator().manual_seed(5)
def rn(*s): return torch.randn(*s,generator=g)
# normalized_entropy convention
lg=rn(1,3,5); mk=torch.tensor([[1.,1.,0.]])
print("norm_ent", R.normalized_entropy(lg,mk))
print("entropy/log5", R.entropy(lg,mk)/torch.log(torch.tensor(5.0)))
# importance_ratio clip None
print("ir None", R.importance_ratio(torch.tensor([0.5]),torch.tensor([0.0]),None))
# reverse_kl_objective
B,T,V=2,3,4
logits=rn(B,T,V); ref=rn(B,T,V); labels=torch.randint(0,V,(B,T),generator=g)
mask=(torch.rand(B,T,generator=g)>0.3).float(); adv=rn(B,T)
print("rkl_obj", R.reverse_kl_objective(logits,ref,labels,mask,adv,0.1))
EOF
python3 /app/_p9.pynorm_ent tensor(0.7432) entropy/log5 tensor(0.7432) ir None tensor([1.6487]) rkl_obj tensor(2.1034) [stdout] norm_ent tensor(0.7432) entropy/log5 tensor(0.7432) ir None tensor([1.6487]) rkl_obj tensor(2.1034)
Solve reverse_kl_objective
cat > /app/_p10.py << 'EOF'
import torch, reforacle as R
g=torch.Generator().manual_seed(5)
def rn(*s): return torch.randn(*s,generator=g)
B,T,V=2,3,4
logits=rn(B,T,V); ref=rn(B,T,V); labels=torch.randint(0,V,(B,T),generator=g)
mask=(torch.rand(B,T,generator=g)>0.3).float(); adv=rn(B,T)
def token_logprobs(lg,lb):
lp=torch.log_softmax(lg,-1)
return lp.gather(-1,lb.unsqueeze(-1)).squeeze(-1)
def mmean(v,m):
return (v*m).sum()/m.sum().clamp(min=1e-9) if m.sum()>0 else torch.tensor(0.)
logp=token_logprobs(logits,labels)
reflp=token_logprobs(ref,labels)
d=logp-reflp
kl=torch.exp(-d)+d-1 # reverse_kl per token
target=R.reverse_kl_objective(logits,ref,labels,mask,adv,0.1)
print("target",target)
for name,cand in [
("-adv*logp+beta*kl", (-adv*logp+0.1*kl)),
("-(adv*logp)-beta*kl", (-adv*logp-0.1*kl)),
("adv*logp+beta*kl", (adv*logp+0.1*kl)),
]:
print(name, mmean(cand,mask).item())
EOF
python3 /app/_p10.pytarget tensor(-0.5832) -adv*logp+beta*kl -0.24233147501945496 -(adv*logp)-beta*kl -0.840238094329834 adv*logp+beta*kl 0.840238094329834 [stdout] target tensor(-0.5832) -adv*logp+beta*kl -0.24233147501945496 -(adv*logp)-beta*kl -0.840238094329834 adv*logp+beta*kl 0.840238094329834
Break down components
cat >> /app/_p10.py << 'EOF'
print("mm(-adv*logp)", mmean(-adv*logp,mask).item())
print("mm(adv*logp)", mmean(adv*logp,mask).item())
print("mm(kl)", mmean(kl,mask).item())
print("mm(logp)", mmean(logp,mask).item())
# maybe -adv*logp - beta*... or +? target -0.5832
# try beta*kl - adv*logp with various
import itertools
base_adv=mmean(-adv*logp,mask).item()
base_kl=mmean(kl,mask).item()
print("check a*base_adv? target/base", -0.5832)
print("mm(-adv*logp)+0.1*mm(kl)", base_adv+0.1*base_kl)
print("mm(-adv*logp)-0.1*mm(kl)", base_adv-0.1*base_kl)
EOF
python3 /app/_p10.pytarget tensor(-0.5832) -adv*logp+beta*kl -0.24233147501945496 -(adv*logp)-beta*kl -0.840238094329834 adv*logp+beta*kl 0.840238094329834 mm(-adv*logp) -0.5412847995758057 mm(adv*logp) 0.5412847995758057 mm(kl) 2.989532947540283 mm(logp) -2.2856521606445312 check a*base_adv? target/base -0.5832 mm(-adv*logp)+0.1*mm(kl) -0.24233150482177734 mm(-adv*logp)-0.1*mm(kl) -0.840238094329834 [stdout] target tensor(-0.5832) -adv*logp+beta*kl -0.24233147501945496 -(adv*logp)-beta*kl -0.840238094329834 adv*logp+beta*kl 0.840238094329834 mm(-adv*logp) -0.5412847995758057 mm(adv*logp) 0.5412847995758057 mm(kl) 2.989532947540283 mm(logp) -2.2856521606445312 check a*base_adv? target/base -0.5832 mm(-adv*logp)+0.1*mm(kl) -0.24233150482177734 mm(-adv*logp)-0.1*mm(kl) -0.840238094329834
More KL candidates
cat >> /app/_p10.py << 'EOF'
fwd=torch.exp(d)-d-1
k1=d; k2=0.5*d*d
print("mm(k1)", mmean(k1,mask).item())
print("mm(fwd)", mmean(fwd,mask).item())
print("mm(k2)", mmean(k2,mask).item())
print("need per-tok mm ~", (-0.5832-(-0.5412847995758057))/0.1)
# maybe adv includes ref: pg with ratio to ref?
ratio=torch.exp(logp-reflp)
print("mm(-adv*ratio)", mmean(-adv*ratio,mask).item())
print("mm(-adv*ratio+0.1kl)", mmean(-adv*ratio+0.1*kl,mask).item())
EOF
python3 /app/_p10.py 2>&1 | tail -6mm(k1) -0.8586515188217163 mm(fwd) 0.9142506122589111 mm(k2) 1.3712430000305176 need per-tok mm ~ -0.4191520042419439 mm(-adv*ratio) 0.45195621252059937 mm(-adv*ratio+0.1kl) 0.7509095668792725 [stdout] mm(k1) -0.8586515188217163 mm(fwd) 0.9142506122589111 mm(k2) 1.3712430000305176 need per-tok mm ~ -0.4191520042419439 mm(-adv*ratio) 0.45195621252059937 mm(-adv*ratio+0.1kl) 0.7509095668792725
Log-ratio candidates
cat >> /app/_p10.py << 'EOF'
print("mm(-adv*d)", mmean(-adv*d,mask).item())
print("mm(-adv*reflp)", mmean(-adv*reflp,mask).item())
print("mm(-adv*d)+0.1*mm(kl)", mmean(-adv*d,mask).item()+0.1*mmean(kl,mask).item())
print("mm(-adv*d + 0.1*kl)", mmean(-adv*d+0.1*kl,mask).item())
print("mm(-adv*d - 0.1*kl)", mmean(-adv*d-0.1*kl,mask).item())
print("mm(adv*d + 0.1*kl)", mmean(adv*d+0.1*kl,mask).item())
# ratio pg with kl subtracted in advantage
print("mm(-(adv-0.1*kl)*logp)", mmean(-(adv-0.1*kl)*logp,mask).item())
EOF
python3 /app/_p10.py 2>&1 | tail -6mm(-adv*reflp) -0.4299217164516449 mm(-adv*d)+0.1*mm(kl) 0.18759019672870636 mm(-adv*d + 0.1*kl) 0.18759021162986755 mm(-adv*d - 0.1*kl) -0.4103164076805115 mm(adv*d + 0.1*kl) 0.4103164076805115 mm(-(adv-0.1*kl)*logp) -1.5442219972610474 [stdout] mm(-adv*reflp) -0.4299217164516449 mm(-adv*d)+0.1*mm(kl) 0.18759019672870636 mm(-adv*d + 0.1*kl) 0.18759021162986755 mm(-adv*d - 0.1*kl) -0.4103164076805115 mm(adv*d + 0.1*kl) 0.4103164076805115 mm(-(adv-0.1*kl)*logp) -1.5442219972610474
Isolate KL and adv terms
cat > /app/_p11.py << 'EOF'
import torch, reforacle as R
g=torch.Generator().manual_seed(5)
def rn(*s): return torch.randn(*s,generator=g)
B,T,V=2,3,4
logits=rn(B,T,V); ref=rn(B,T,V); labels=torch.randint(0,V,(B,T),generator=g)
mask=torch.ones(B,T); adv=rn(B,T)
def tlp(lg,lb):
return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
logp=tlp(logits,labels); reflp=tlp(ref,labels); d=logp-reflp
kl=torch.exp(-d)+d-1
def mm(v,m): return (v*m).sum()/m.sum()
print("adv=0:", R.reverse_kl_objective(logits,ref,labels,mask,torch.zeros(B,T),0.1).item())
print(" 0.1*mm(kl)=",0.1*mm(kl,mask).item())
print(" 0.1*mm(k1)=",0.1*mm(d,mask).item())
print("beta=0:", R.reverse_kl_objective(logits,ref,labels,mask,adv,0.0).item())
print(" mm(-adv*logp)=",mm(-adv*logp,mask).item())
print(" mm(-adv*d)=",mm(-adv*d,mask).item())
print(" mm(adv*logp)=",mm(adv*logp,mask).item())
EOF
python3 /app/_p11.pyadv=0: 0.7638916969299316 0.1*mm(kl)= 0.3819458961486817 0.1*mm(k1)= -0.08011066913604736 beta=0: 0.21610014140605927 mm(-adv*logp)= -0.39875224232673645 mm(-adv*d)= -0.7132778763771057 mm(adv*logp)= 0.39875224232673645 [stdout] adv=0: 0.7638916969299316 0.1*mm(kl)= 0.3819458961486817 0.1*mm(k1)= -0.08011066913604736 beta=0: 0.21610014140605927 mm(-adv*logp)= -0.39875224232673645 mm(-adv*d)= -0.7132778763771057 mm(adv*logp)= 0.39875224232673645
True KL and ratio
cat >> /app/_p11.py << 'EOF'
lpp=torch.log_softmax(logits,-1); lpr=torch.log_softmax(ref,-1)
pp=lpp.exp()
trueKL=(pp*(lpp-lpr)).sum(-1) # reverse KL(policy||ref) per token
print("--- true distributional ---")
print("0.1*mm(trueKL)=",0.1*mm(trueKL,mask).item(), " x2=",0.2*mm(trueKL,mask).item())
ratio=torch.exp(d)
print("mm(-adv*ratio)=",mm(-adv*ratio,mask).item())
print("mm(adv*ratio)=",mm(adv*ratio,mask).item())
print("mm(-adv*logp)+? target beta0=0.2161")
# maybe advantage term is -mm(adv*logp) but adv whitened? or masked_mean differ
EOF
python3 /app/_p11.py 2>&1 | tail -6mm(adv*logp)= 0.39875224232673645 --- true distributional --- 0.1*mm(trueKL)= 0.08726447224617005 x2= 0.1745289444923401 mm(-adv*ratio)= -0.8311963081359863 mm(adv*ratio)= 0.8311963081359863 mm(-adv*logp)+? target beta0=0.2161 [stdout] mm(adv*logp)= 0.39875224232673645 --- true distributional --- 0.1*mm(trueKL)= 0.08726447224617005 x2= 0.1745289444923401 mm(-adv*ratio)= -0.8311963081359863 mm(adv*ratio)= 0.8311963081359863 mm(-adv*logp)+? target beta0=0.2161
Exact decomposition
cat >> /app/_p11.py << 'EOF'
print("=== exact ===")
o0=R.reverse_kl_objective(logits,ref,labels,mask,torch.zeros(B,T),0.1).item()
print("obj adv0 beta0.1", o0, " /0.1=", o0/0.1)
print("sum(kl)",kl.sum().item(),"mean",kl.mean().item())
print("sum(trueKL)",trueKL.sum().item(),"mean",trueKL.mean().item())
# beta=0 case decompose
ob=R.reverse_kl_objective(logits,ref,labels,mask,adv,0.0).item()
print("obj beta0", ob)
print("mm(-adv*logp)",mm(-adv*logp,mask).item())
print("mm(-adv*reflp)",mm(-adv*reflp,mask).item())
# maybe adv gets whitened inside
aw=(adv-adv.mean())/adv.std()
print("mm(-aw*logp)",mm(-aw*logp,mask).item())
EOF
python3 /app/_p11.py 2>&1 | tail -9mm(-adv*logp)+? target beta0=0.2161 === exact === obj adv0 beta0.1 0.7638916969299316 /0.1= 7.638916969299316 sum(kl) 22.9167537689209 mean 3.8194589614868164 sum(trueKL) 5.235868453979492 mean 0.8726447224617004 obj beta0 0.21610014140605927 mm(-adv*logp) -0.39875224232673645 mm(-adv*reflp) 0.31452563405036926 mm(-aw*logp) -0.6378178000450134 [stdout] mm(-adv*logp)+? target beta0=0.2161 === exact === obj adv0 beta0.1 0.7638916969299316 /0.1= 7.638916969299316 sum(kl) 22.9167537689209 mean 3.8194589614868164 sum(trueKL) 5.235868453979492 mean 0.8726447224617004 obj beta0 0.21610014140605927 mm(-adv*logp) -0.39875224232673645 mm(-adv*reflp) 0.31452563405036926 mm(-aw*logp) -0.6378178000450134
Estimator means
cat >> /app/_p11.py << 'EOF'
print("=== estimator means (mask all ones) ===")
for nm,val in [("k1",d),("k2",0.5*d*d),("k3rev",kl),("fwd",torch.exp(d)-d-1),
("sym",torch.cosh(d)-1),("abs_d",d.abs()),("d2",d*d)]:
print(nm, mm(val,mask).item())
print("need 7.639")
EOF
python3 /app/_p11.py 2>&1 | tail -9=== estimator means (mask all ones) === k1 -0.8011066913604736 k2 1.6306487321853638 k3rev 3.8194589614868164 fwd 1.1105259656906128 sym 2.4649922847747803 abs_d 1.612099289894104 d2 3.2612974643707275 need 7.639 [stdout] === estimator means (mask all ones) === k1 -0.8011066913604736 k2 1.6306487321853638 k3rev 3.8194589614868164 fwd 1.1105259656906128 sym 2.4649922847747803 abs_d 1.612099289894104 d2 3.2612974643707275 need 7.639
Reduction hypothesis test
cat >> /app/_p11.py << 'EOF'
print("=== reduction test ===")
def mm_ax1(v,m): return (v*m).sum(1)/m.sum(1)
kterm=mm_ax1(kl,mask)
print("per-seq kl mean", kterm, "sum", kterm.sum().item())
aterm=mm_ax1(-adv*logp,mask)
print("per-seq adv mean", aterm, "sum", aterm.sum().item(), "mean", aterm.mean().item())
print("obj beta0 target 0.21610")
# combined full
full=R.reverse_kl_objective(logits,ref,labels,mask,adv,0.1).item()
print("full target", full)
print("cand sum: ", (mm_ax1(-adv*logp,mask)+0.1*mm_ax1(kl,mask)).sum().item())
print("cand mean:", (mm_ax1(-adv*logp,mask)+0.1*mm_ax1(kl,mask)).mean().item())
EOF
python3 /app/_p11.py 2>&1 | tail -7=== reduction test === per-seq kl mean tensor([7.2879, 0.3511]) sum 7.638917446136475 per-seq adv mean tensor([-1.4003, 0.6028]) sum -0.7975046038627625 mean -0.3987523019313812 obj beta0 target 0.21610 full target 0.9799919128417969 cand sum: -0.033612847328186035 cand mean: -0.016806423664093018 [stdout] === reduction test === per-seq kl mean tensor([7.2879, 0.3511]) sum 7.638917446136475 per-seq adv mean tensor([-1.4003, 0.6028]) sum -0.7975046038627625 mean -0.3987523019313812 obj beta0 target 0.21610 full target 0.9799919128417969 cand sum: -0.033612847328186035 cand mean: -0.016806423664093018
Brute force beta0 term
cat >> /app/_p11.py << 'EOF'
print("=== beta0 brute (target 0.21610) ===")
ratio=torch.exp(d)
cands={
"mean(-adv*logp)":(-adv*logp),
"mean(adv*logp)":(adv*logp),
"mean(-adv*d)":(-adv*d),
"mean(-adv*ratio)":(-adv*ratio),
"mean(adv*ratio)":(adv*ratio),
"mean(-adv*reflp)":(-adv*reflp),
"mean(-(adv)*(logp-reflp.detach))":(-adv*d),
}
for n,v in cands.items(): print(n, mm(v,mask).item(), " sum_ax1", mm_ax1(v,mask).sum().item())
EOF
python3 /app/_p11.py 2>&1 | tail -8=== beta0 brute (target 0.21610) === mean(-adv*logp) -0.39875224232673645 sum_ax1 -0.7975046038627625 mean(adv*logp) 0.39875224232673645 sum_ax1 0.7975046038627625 mean(-adv*d) -0.7132778763771057 sum_ax1 -1.426555871963501 mean(-adv*ratio) -0.8311963081359863 sum_ax1 -1.6623926162719727 mean(adv*ratio) 0.8311963081359863 sum_ax1 1.6623926162719727 mean(-adv*reflp) 0.31452563405036926 sum_ax1 0.6290512681007385 mean(-(adv)*(logp-reflp.detach)) -0.7132778763771057 sum_ax1 -1.426555871963501 [stdout] === beta0 brute (target 0.21610) === mean(-adv*logp) -0.39875224232673645 sum_ax1 -0.7975046038627625 mean(adv*logp) 0.39875224232673645 sum_ax1 0.7975046038627625 mean(-adv*d) -0.7132778763771057 sum_ax1 -1.426555871963501 mean(-adv*ratio) -0.8311963081359863 sum_ax1 -1.6623926162719727 mean(adv*ratio) 0.8311963081359863 sum_ax1 1.6623926162719727 mean(-adv*reflp) 0.31452563405036926 sum_ax1 0.6290512681007385 mean(-(adv)*(logp-reflp.detach)) -0.7132778763771057 sum_ax1 -1.426555871963501
Whitened advantage term
cat >> /app/_p11.py << 'EOF'
print("=== whitened adv (target A=0.21610) ===")
am=adv.mean(); asb=adv.std(unbiased=False); asu=adv.std(unbiased=True)
for nm,aa in [("biased",(adv-am)/asb),("unbiased",(adv-am)/asu),("noshift_b",adv/asb)]:
print(nm,"mean(-a*logp)",mm(-aa*logp,mask).item()," mean(a*logp)",mm(aa*logp,mask).item(),
" sumax1(-a*logp)",mm_ax1(-aa*logp,mask).sum().item())
EOF
python3 /app/_p11.py 2>&1 | tail -4=== whitened adv (target A=0.21610) === biased mean(-a*logp) -0.6986944079399109 mean(a*logp) 0.6986944079399109 sumax1(-a*logp) -1.3973889350891113 unbiased mean(-a*logp) -0.6378178000450134 mean(a*logp) 0.6378178000450134 sumax1(-a*logp) -1.2756356000900269 noshift_b mean(-a*logp) -0.554550051689148 mean(a*logp) 0.554550051689148 sumax1(-a*logp) -1.109100103378296 [stdout] === whitened adv (target A=0.21610) === biased mean(-a*logp) -0.6986944079399109 mean(a*logp) 0.6986944079399109 sumax1(-a*logp) -1.3973889350891113 unbiased mean(-a*logp) -0.6378178000450134 mean(a*logp) 0.6378178000450134 sumax1(-a*logp) -1.2756356000900269 noshift_b mean(-a*logp) -0.554550051689148 mean(a*logp) 0.554550051689148 sumax1(-a*logp) -1.109100103378296
Per-sequence A
cat >> /app/_p11.py << 'EOF'
print("=== per-seq A decomposition ===")
adv0=adv.clone(); adv0[1]=0
adv1=adv.clone(); adv1[0]=0
A0=R.reverse_kl_objective(logits,ref,labels,mask,adv0,0.0).item()
A1=R.reverse_kl_objective(logits,ref,labels,mask,adv1,0.0).item()
print("A0",A0,"A1",A1,"sum",A0+A1,"target A",0.21610014)
print("row0 -adv*logp mean", (-adv[0]*logp[0]).mean().item(), " row1", (-adv[1]*logp[1]).mean().item())
print("row0 adv", adv[0].tolist(), "logp0", logp[0].tolist())
EOF
python3 /app/_p11.py 2>&1 | tail -5noshift_b mean(-a*logp) -0.554550051689148 mean(a*logp) 0.554550051689148 sumax1(-a*logp) -1.109100103378296 === per-seq A decomposition === A0 -0.9549288749694824 A1 1.1710290908813477 sum 0.21610021591186523 target A 0.21610014 row0 -adv*logp mean -1.4003394842147827 row1 0.6028348803520203 row0 adv [0.31667575240135193, -1.1957919597625732, -0.31266358494758606] logp0 [-3.8347020149230957, -3.6795475482940674, -3.2475860118865967] [stdout] noshift_b mean(-a*logp) -0.554550051689148 mean(a*logp) 0.554550051689148 sumax1(-a*logp) -1.109100103378296 === per-seq A decomposition === A0 -0.9549288749694824 A1 1.1710290908813477 sum 0.21610021591186523 target A 0.21610014 row0 -adv*logp mean -1.4003394842147827 row1 0.6028348803520203 row0 adv [0.31667575240135193, -1.1957919597625732, -0.31266358494758606] logp0 [-3.8347020149230957, -3.6795475482940674, -3.2475860118865967]
Row0 detailed candidates
cat >> /app/_p11.py << 'EOF'
print("=== row0 detail, A0=-0.95493 ===")
r0=torch.exp(d[0]);
print("reflp0", reflp[0].tolist(), "ratio0", r0.tolist(), "d0", d[0].tolist())
a0=adv[0]
for nm,v in [("-a*logp",-a0*logp[0]),("-a*ratio",-a0*r0),("a*ratio",a0*r0),
("-a*d",-a0*d[0]),("a*d",a0*d[0]),("-a*reflp",-a0*reflp[0])]:
print(nm,"mean",v.mean().item(),"sum",v.sum().item())
EOF
python3 /app/_p11.py 2>&1 | tail -7reflp0 [-1.1920806169509888, -1.5303988456726074, -1.0226691961288452] ratio0 [0.07117444276809692, 0.11658336222171783, 0.10807639360427856] d0 [-2.6426215171813965, -2.14914870262146, -2.224916934967041] -a*logp mean -1.4003394842147827 sum -4.201018333435059 -a*ratio mean 0.05022059381008148 sum 0.15066178143024445 a*ratio mean -0.05022059381008148 sum -0.15066178143024445 -a*d mean -0.8095770478248596 sum -2.4287312030792236 a*d mean 0.8095770478248596 sum 2.4287312030792236 -a*reflp mean -0.5907623171806335 sum -1.7722870111465454 [stdout] reflp0 [-1.1920806169509888, -1.5303988456726074, -1.0226691961288452] ratio0 [0.07117444276809692, 0.11658336222171783, 0.10807639360427856] d0 [-2.6426215171813965, -2.14914870262146, -2.224916934967041] -a*logp mean -1.4003394842147827 sum -4.201018333435059 -a*ratio mean 0.05022059381008148 sum 0.15066178143024445 a*ratio mean -0.05022059381008148 sum -0.15066178143024445 -a*d mean -0.8095770478248596 sum -2.4287312030792236 a*d mean 0.8095770478248596 sum 2.4287312030792236 -a*reflp mean -0.5907623171806335 sum -1.7722870111465454
Linearity and ref dependence
cat >> /app/_p11.py << 'EOF'
print("=== structure probes ===")
print("A(adv)", R.reverse_kl_objective(logits,ref,labels,mask,adv,0.0).item())
print("A(2*adv)", R.reverse_kl_objective(logits,ref,labels,mask,2*adv,0.0).item())
print("A(adv), ref=logits", R.reverse_kl_objective(logits,logits,labels,mask,adv,0.0).item())
print("A(adv,beta.1) ref=logits", R.reverse_kl_objective(logits,logits,labels,mask,adv,0.1).item())
EOF
python3 /app/_p11.py 2>&1 | tail -5=== structure probes === A(adv) 0.21610014140605927 A(2*adv) 0.43220028281211853 A(adv), ref=logits 0.21610014140605927 A(adv,beta.1) ref=logits 0.21610014140605927 [stdout] === structure probes === A(adv) 0.21610014140605927 A(2*adv) 0.43220028281211853 A(adv), ref=logits 0.21610014140605927 A(adv,beta.1) ref=logits 0.21610014140605927
Recover linear functional weights
cat >> /app/_p11.py << 'EOF'
print("=== recover linear weights w ===")
W=torch.zeros(B,T)
for i in range(B):
for t in range(T):
e=torch.zeros(B,T); e[i,t]=1.0
W[i,t]=R.reverse_kl_objective(logits,ref,labels,mask,e,0.0).item()
print("W\n",W)
print("logp\n",logp)
print("W/logp\n", W/logp)
print("-logp/6\n", -logp/6)
EOF
python3 /app/_p11.py 2>&1 | tail -12W
tensor([[0.8194, 0.8247, 0.7302],
[0.8194, 0.8247, 0.7302]])
logp
tensor([[-3.8347, -3.6795, -3.2476],
[-1.0818, -1.2684, -1.1337]])
W/logp
tensor([[-0.2137, -0.2241, -0.2248],
[-0.7574, -0.6502, -0.6441]])
-logp/6
tensor([[0.6391, 0.6133, 0.5413],
[0.1803, 0.2114, 0.1889]])
[stdout]
W
tensor([[0.8194, 0.8247, 0.7302],
[0.8194, 0.8247, 0.7302]])
logp
tensor([[-3.8347, -3.6795, -3.2476],
[-1.0818, -1.2684, -1.1337]])
W/logp
tensor([[-0.2137, -0.2241, -0.2248],
[-0.7574, -0.6502, -0.6441]])
-logp/6
tensor([[0.6391, 0.6133, 0.5413],
[0.1803, 0.2114, 0.1889]])Test W logit dependence
cat >> /app/_p11.py << 'EOF'
print("=== does W depend on logits? ===")
g2=torch.Generator().manual_seed(99)
logits2=torch.randn(B,T,V,generator=g2)*3
W2=torch.zeros(B,T)
for i in range(B):
for t in range(T):
e=torch.zeros(B,T); e[i,t]=1.0
W2[i,t]=R.reverse_kl_objective(logits2,ref,labels,mask,e,0.0).item()
print("W2\n",W2)
logp2=torch.log_softmax(logits2,-1).gather(-1,labels.unsqueeze(-1)).squeeze(-1)
print("logp2\n",logp2)
print("exp(logp2)\n", logp2.exp())
EOF
python3 /app/_p11.py 2>&1 | tail -8tensor([[1.9100, 2.4247, 0.5501],
[1.9100, 2.4247, 0.5501]])
logp2
tensor([[ -0.9917, -8.7101, -0.1141],
[-10.4685, -5.8379, -3.1863]])
exp(logp2)
tensor([[3.7096e-01, 1.6490e-04, 8.9214e-01],
[2.8418e-05, 2.9148e-03, 4.1325e-02]])
[stdout]
tensor([[1.9100, 2.4247, 0.5501],
[1.9100, 2.4247, 0.5501]])
logp2
tensor([[ -0.9917, -8.7101, -0.1141],
[-10.4685, -5.8379, -3.1863]])
exp(logp2)
tensor([[3.7096e-01, 1.6490e-04, 8.9214e-01],
[2.8418e-05, 2.9148e-03, 4.1325e-02]])Pin reverse_kl_objective formula B=3
cat > /app/_p12.py << 'EOF'
import torch, reforacle as R
g=torch.Generator().manual_seed(7)
B,T,V=3,2,4
logits=torch.randn(B,T,V,generator=g); ref=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g); mask=torch.ones(B,T)
logp=torch.log_softmax(logits,-1).gather(-1,labels.unsqueeze(-1)).squeeze(-1)
W=torch.zeros(B,T)
for i in range(B):
for t in range(T):
e=torch.zeros(B,T); e[i,t]=1.0
W[i,t]=R.reverse_kl_objective(logits,ref,labels,mask,e,0.0).item()
print("W\n",W)
N=B*T
print("colsum logp /N: col-wise -(sum_i logp)/N")
print(-(logp.sum(0,keepdim=True)/N).expand(B,T))
print("-logp/N\n", -logp/N)
print("-logp/B\n", -logp/B)
EOF
python3 /app/_p12.pyW
tensor([[0.7046, 0.5989],
[0.7046, 0.5989],
[0.7046, 0.5989]])
colsum logp /N: col-wise -(sum_i logp)/N
tensor([[0.7046, 0.5989],
[0.7046, 0.5989],
[0.7046, 0.5989]])
-logp/N
tensor([[0.4869, 0.1417],
[0.0587, 0.1134],
[0.1590, 0.3437]])
-logp/B
tensor([[0.9738, 0.2834],
[0.1173, 0.2269],
[0.3180, 0.6875]])
[stdout]
W
tensor([[0.7046, 0.5989],
[0.7046, 0.5989],
[0.7046, 0.5989]])
colsum logp /N: col-wise -(sum_i logp)/N
tensor([[0.7046, 0.5989],
[0.7046, 0.5989],
[0.7046, 0.5989]])
-logp/N
tensor([[0.4869, 0.1417],
[0.0587, 0.1134],
[0.1590, 0.3437]])
-logp/B
tensor([[0.9738, 0.2834],
[0.1173, 0.2269],
[0.3180, 0.6875]])Mask handling for reverse_kl_objective
cat > /app/_p13.py << 'EOF'
import torch, reforacle as R
g=torch.Generator().manual_seed(7)
B,T,V=3,2,4
logits=torch.randn(B,T,V,generator=g); ref=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g)
mask=torch.tensor([[1.,1.],[1.,0.],[0.,1.]])
logp=torch.log_softmax(logits,-1).gather(-1,labels.unsqueeze(-1)).squeeze(-1)
W=torch.zeros(B,T)
for i in range(B):
for t in range(T):
e=torch.zeros(B,T); e[i,t]=1.0
W[i,t]=R.reverse_kl_objective(logits,ref,labels,mask,e,0.0).item()
print("W (masked)\n",W)
Nsum=mask.sum().item()
# hypo: W_it = -(sum_j logp_jt*m_jt)/Nsum * m_it
colsum=(logp*mask).sum(0,keepdim=True)
print("cand -(colsum_logp)/sum(mask) * m\n", (-(colsum)/Nsum)*mask )
print("cand with per-col denom colmask")
colmask=mask.sum(0,keepdim=True)
print((-(colsum)/colmask.clamp(min=1))*mask)
EOF
python3 /app/_p13.pyW (masked)
tensor([[0.8183, 0.7281],
[0.8183, 0.7281],
[0.8183, 0.7281]])
cand -(colsum_logp)/sum(mask) * m
tensor([[0.8183, 0.7281],
[0.8183, 0.0000],
[0.0000, 0.7281]])
cand with per-col denom colmask
tensor([[1.6367, 1.4563],
[1.6367, 0.0000],
[0.0000, 1.4563]])
[stdout]
W (masked)
tensor([[0.8183, 0.7281],
[0.8183, 0.7281],
[0.8183, 0.7281]])
cand -(colsum_logp)/sum(mask) * m
tensor([[0.8183, 0.7281],
[0.8183, 0.0000],
[0.0000, 0.7281]])
cand with per-col denom colmask
tensor([[1.6367, 1.4563],
[1.6367, 0.0000],
[0.0000, 1.4563]])Test full reverse_kl_objective candidate
cat > /app/_p14.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
def cand(logits,ref,labels,mask,adv,beta):
logp=tlp(logits,labels); reflp=tlp(ref,labels); d=logp-reflp
kl=torch.exp(-d)+d-1
M=mask.sum()
term_adv=-(( (logp*mask).sum(0) )*adv.sum(0)).sum()/M
# KL: sum over batch of masked_mean over tokens
seqden=mask.sum(1).clamp(min=1e-9)
term_kl=beta*((kl*mask).sum(1)/seqden).sum()
return term_adv+term_kl
for seed in [1,2,3,4]:
g=torch.Generator().manual_seed(seed)
B,T,V=torch.randint(2,5,(1,),generator=g).item(),torch.randint(2,5,(1,),generator=g).item(),5
logits=torch.randn(B,T,V,generator=g);ref=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g)
mask=(torch.rand(B,T,generator=g)>0.3).float()
adv=torch.randn(B,T,generator=g); beta=0.1
r=R.reverse_kl_objective(logits,ref,labels,mask,adv,beta).item()
c=cand(logits,ref,labels,mask,adv,beta).item()
print(seed,"ref",round(r,5),"cand",round(c,5),"match",abs(r-c)<1e-4)
EOF
python3 /app/_p14.py1 ref 1.41706 cand 1.45427 match False 2 ref 0.86723 cand 0.80732 match False 3 ref -0.43945 cand -0.41503 match False 4 ref -0.5213 cand -0.55634 match False [stdout] 1 ref 1.41706 cand 1.45427 match False 2 ref 0.86723 cand 0.80732 match False 3 ref -0.43945 cand -0.41503 match False 4 ref -0.5213 cand -0.55634 match False
Separate adv and kl term mask handling
cat > /app/_p15.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
for seed in [1,2,3]:
g=torch.Generator().manual_seed(seed)
B,T,V=3,3,5
logits=torch.randn(B,T,V,generator=g);ref=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g)
mask=(torch.rand(B,T,generator=g)>0.3).float()
adv=torch.randn(B,T,generator=g)
logp=tlp(logits,labels); reflp=tlp(ref,labels); d=logp-reflp
kl=torch.exp(-d)+d-1
M=mask.sum()
# adv term (beta 0)
at=R.reverse_kl_objective(logits,ref,labels,mask,adv,0.0).item()
c1=-(((logp*mask).sum(0))*adv.sum(0)).sum()/M
c2=-((logp*mask)*adv).sum()/M # standard masked mean adv*logp
c3=-((logp*mask)*adv*mask).sum()/M
print(seed,"ADV ref",round(at,5),"colsum",round(c1.item(),5),"std_mm",round(c2.item(),5))
# kl term (adv 0)
kt=R.reverse_kl_objective(logits,ref,labels,mask,torch.zeros_like(adv),0.1).item()
seqden=mask.sum(1).clamp(min=1e-9)
k1=0.1*((kl*mask).sum(1)/seqden).sum()
k2=0.1*((kl*mask).sum()/M) # global masked mean
k3=0.1*(kl*mask).sum() # masked sum
print(" KL ref",round(kt,5),"seqmean_sum",round(k1.item(),5),"glob_mm",round(k2.item(),5),"msum",round(k3.item(),5))
EOF
python3 /app/_p15.py1 ADV ref 0.44982 colsum 0.44982 std_mm 0.37927 KL ref 0.20008 seqmean_sum 0.19831 glob_mm 0.06669 msum 0.46684 2 ADV ref 2.46023 colsum 2.46023 std_mm 0.66411 KL ref 0.29037 seqmean_sum 0.54263 glob_mm 0.09679 msum 0.58073 3 ADV ref -3.13367 colsum -3.13367 std_mm -0.49742 KL ref 0.40452 seqmean_sum 0.4237 glob_mm 0.13484 msum 0.94389 [stdout] 1 ADV ref 0.44982 colsum 0.44982 std_mm 0.37927 KL ref 0.20008 seqmean_sum 0.19831 glob_mm 0.06669 msum 0.46684 2 ADV ref 2.46023 colsum 2.46023 std_mm 0.66411 KL ref 0.29037 seqmean_sum 0.54263 glob_mm 0.09679 msum 0.58073 3 ADV ref -3.13367 colsum -3.13367 std_mm -0.49742 KL ref 0.40452 seqmean_sum 0.4237 glob_mm 0.13484 msum 0.94389
Recover KL reduction coefficients
cat > /app/_p16.py << 'EOF'
import torch, reforacle as R
B,T,V=3,3,5
labels=torch.zeros(B,T,dtype=torch.long)
mask=torch.tensor([[1.,1.,0.],[1.,0.,1.],[1.,1.,1.]])
ref=torch.zeros(B,T,V) # uniform, reflp=-log5
refl=-torch.log(torch.tensor(5.0))
C=torch.zeros(B,T)
for i in range(B):
for t in range(T):
logits=torch.zeros(B,T,V)
logits[i,t,0]=2.0 # change logp at this token
lp=torch.log_softmax(logits[i,t],-1)[0]
d=lp-refl
kl=torch.exp(-d)+d-1
val=R.reverse_kl_objective(logits,ref,labels,mask,torch.zeros(B,T),1.0).item()
C[i,t]=val/kl.item() if kl.item()!=0 else 0
print("C (kl coeffs), mask=\n",mask)
print(C)
seqden=mask.sum(1,keepdim=True)
print("mask/seqden\n", mask/seqden)
print("mask/summask\n", mask/mask.sum())
EOF
python3 /app/_p16.pyC (kl coeffs), mask=
tensor([[1., 1., 0.],
[1., 0., 1.],
[1., 1., 1.]])
tensor([[0.4286, 0.4286, 0.0000],
[0.4286, 0.0000, 0.4286],
[0.4286, 0.4286, 0.4286]])
mask/seqden
tensor([[0.5000, 0.5000, 0.0000],
[0.5000, 0.0000, 0.5000],
[0.3333, 0.3333, 0.3333]])
mask/summask
tensor([[0.1429, 0.1429, 0.0000],
[0.1429, 0.0000, 0.1429],
[0.1429, 0.1429, 0.1429]])
[stdout]
C (kl coeffs), mask=
tensor([[1., 1., 0.],
[1., 0., 1.],
[1., 1., 1.]])
tensor([[0.4286, 0.4286, 0.0000],
[0.4286, 0.0000, 0.4286],
[0.4286, 0.4286, 0.4286]])
mask/seqden
tensor([[0.5000, 0.5000, 0.0000],
[0.5000, 0.0000, 0.5000],
[0.3333, 0.3333, 0.3333]])
mask/summask
tensor([[0.1429, 0.1429, 0.0000],
[0.1429, 0.0000, 0.1429],
[0.1429, 0.1429, 0.1429]])Careful KL term isolation
cat > /app/_p17.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
g=torch.Generator().manual_seed(1)
B,T,V=3,3,5
logits=torch.randn(B,T,V,generator=g);ref=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g)
mask=(torch.rand(B,T,generator=g)>0.3).float()
logp=tlp(logits,labels); reflp=tlp(ref,labels); d=logp-reflp
kl=torch.exp(-d)+d-1
kt=R.reverse_kl_objective(logits,ref,labels,mask,torch.zeros(B,T),1.0).item()
print("mask\n",mask)
print("kl\n",kl)
print("ref kt(beta1)",kt)
print("masked_sum kl", (kl*mask).sum().item())
print("masked_mean kl", ((kl*mask).sum()/mask.sum()).item())
print("sum over batch of masked_mean_t:", ((kl*mask).sum(1)/mask.sum(1).clamp(min=1)).sum().item())
# maybe mean over batch of masked_sum_t
print("colsum coupling like adv? -(masksum_kl col)*ones", None)
EOF
python3 /app/_p17.pymask
tensor([[1., 1., 1.],
[1., 1., 1.],
[1., 0., 0.]])
kl
tensor([[1.0176, 0.2651, 0.2840],
[0.3212, 1.3013, 0.8388],
[0.6404, 0.1944, 1.6581]])
ref kt(beta1) 2.000753164291382
masked_sum kl 4.668424129486084
masked_mean kl 0.6669177412986755
sum over batch of masked_mean_t: 1.9830830097198486
colsum coupling like adv? -(masksum_kl col)*ones None
[stdout]
mask
tensor([[1., 1., 1.],
[1., 1., 1.],
[1., 0., 0.]])
kl
tensor([[1.0176, 0.2651, 0.2840],
[0.3212, 1.3013, 0.8388],
[0.6404, 0.1944, 1.6581]])
ref kt(beta1) 2.000753164291382
masked_sum kl 4.668424129486084
masked_mean kl 0.6669177412986755
sum over batch of masked_mean_t: 1.9830830097198486
colsum coupling like adv? -(masksum_kl col)*ones NoneTest reverse_kl_objective full hypothesis
cat > /app/_p18.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
def cand(logits,ref,labels,mask,adv,beta):
logp=tlp(logits,labels); reflp=tlp(ref,labels); d=logp-reflp
kl=torch.exp(-d)+d-1
B=logits.shape[0]; M=mask.sum()
term_adv=-(((logp*mask).sum(0))*adv.sum(0)).sum()/M
term_kl=beta*B*((kl*mask).sum()/M)
return term_adv+term_kl
for seed in range(1,7):
g=torch.Generator().manual_seed(seed)
B=torch.randint(2,5,(1,),generator=g).item();T=torch.randint(2,5,(1,),generator=g).item();V=6
logits=torch.randn(B,T,V,generator=g);ref=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g)
mask=(torch.rand(B,T,generator=g)>0.3).float()
adv=torch.randn(B,T,generator=g); beta=torch.rand(1,generator=g).item()
r=R.reverse_kl_objective(logits,ref,labels,mask,adv,beta).item()
c=cand(logits,ref,labels,mask,adv,beta).item()
print(seed,round(r,5),round(c,5),abs(r-c)<1e-4)
EOF
python3 /app/_p18.py1 -5.35465 -5.35465 True 2 0.06041 0.06041 True 3 2.90422 2.90422 True 4 -1.26715 -1.26715 True 5 -2.28249 -2.28249 True 6 7.86388 7.86388 True [stdout] 1 -5.35465 -5.35465 True 2 0.06041 0.06041 True 3 2.90422 2.90422 True 4 -1.26715 -1.26715 True 5 -2.28249 -2.28249 True 6 7.86388 7.86388 True
Probe reductions, shapes, edge cases
cat > /app/_p19.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
g=torch.Generator().manual_seed(3)
# entropy reduction B=2
logits=torch.randn(2,3,4,generator=g); mask=torch.ones(2,3)
p=logits.softmax(-1); ent=-(p*logits.log_softmax(-1)).sum(-1)
print("entropy ref", R.entropy(logits,mask).item(), "plain mean", ent.mean().item(), "B*mean", 2*ent.mean().item())
# grpo eps with zero-variance group
r=torch.tensor([5.,5.,5.,1.,2.,3.])
print("grpo constgroup scaleT", R.grpo_advantages(r,3,True))
# gae 2d?
try:
rw=torch.randn(2,3,generator=g); vv=torch.randn(2,3,generator=g); nv=torch.randn(2,generator=g)
print("gae2d", R.gae(rw,vv,nv,0.9,0.95).shape)
except Exception as e: print("gae2d ERR",e)
# discounted_returns 2d?
try:
print("discret2d", R.discounted_returns(torch.randn(2,3,generator=g),0.9).shape)
except Exception as e: print("dr2d ERR",e)
# cross_entropy all-ignored
print("ce all ignore", R.cross_entropy(torch.randn(2,4,generator=g),torch.tensor([-100,-100]),-100))
# top_k ties
print("top_k tie", R.top_k_mask(torch.tensor([[1.,1.,1.,0.]]),2))
EOF
python3 /app/_p19.pyentropy ref 1.22975754737854 plain mean 1.22975754737854 B*mean 2.45951509475708 grpo constgroup scaleT tensor([ 0.0000, 0.0000, 0.0000, -1.2247, 0.0000, 1.2247]) gae2d ERR The size of tensor a (3) must match the size of tensor b (2) at non-singleton dimension 0 discret2d torch.Size([2, 3]) ce all ignore tensor(0.) top_k tie tensor([[ True, True, True, False]]) [stdout] entropy ref 1.22975754737854 plain mean 1.22975754737854 B*mean 2.45951509475708 grpo constgroup scaleT tensor([ 0.0000, 0.0000, 0.0000, -1.2247, 0.0000, 1.2247]) gae2d ERR The size of tensor a (3) must match the size of tensor b (2) at non-singleton dimension 0 discret2d torch.Size([2, 3]) ce all ignore tensor(0.) top_k tie tensor([[ True, True, True, False]])
grpo eps and discounted_returns axis
cat > /app/_p20.py << 'EOF'
import torch, reforacle as R
# grpo eps: small std group
r=torch.tensor([0.,0.0001,0.,0.,0.,0.0]) # group1 tiny std
out=R.grpo_advantages(r,3,True)
print("grpo tiny", out)
g1=r[:3]; c=g1-g1.mean(); sb=g1.std(unbiased=False)
for eps in [0,1e-8,1e-6,1e-5,1e-4]:
print("eps",eps,(c/(sb+eps)).tolist())
# discounted_returns axis for 2d
x=torch.tensor([[1.,2.,3.],[10.,20.,30.]])
print("dr2d\n", R.discounted_returns(x,0.5))
EOF
python3 /app/_p20.pygrpo tiny tensor([-0.6924, 1.3848, -0.6924, 0.0000, 0.0000, 0.0000])
eps 0 [-0.7071068286895752, 1.4142135381698608, -0.7071068286895752]
eps 1e-08 [-0.7069568634033203, 1.413913607597351, -0.7069568634033203]
eps 1e-06 [-0.6924183964729309, 1.3848366737365723, -0.6924183964729309]
eps 1e-05 [-0.5833579301834106, 1.1667157411575317, -0.5833579301834106]
eps 0.0001 [-0.2265409380197525, 0.4530818462371826, -0.2265409380197525]
dr2d
tensor([[ 6., 12., 18.],
[10., 20., 30.]])
[stdout]
grpo tiny tensor([-0.6924, 1.3848, -0.6924, 0.0000, 0.0000, 0.0000])
eps 0 [-0.7071068286895752, 1.4142135381698608, -0.7071068286895752]
eps 1e-08 [-0.7069568634033203, 1.413913607597351, -0.7069568634033203]
eps 1e-06 [-0.6924183964729309, 1.3848366737365723, -0.6924183964729309]
eps 1e-05 [-0.5833579301834106, 1.1667157411575317, -0.5833579301834106]
eps 0.0001 [-0.2265409380197525, 0.4530818462371826, -0.2265409380197525]
dr2d
tensor([[ 6., 12., 18.],
[10., 20., 30.]])Probe dpo_sequence_loss
cat > /app/_p21.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
def seqlp(lg,lb,m,norm):
t=(tlp(lg,lb)*m).sum(-1)
return t/m.sum(-1) if norm else t
g=torch.Generator().manual_seed(11)
B,T,V=2,4,5
pc_l=torch.randn(B,T,V,generator=g);pr_l=torch.randn(B,T,V,generator=g)
rc_l=torch.randn(B,T,V,generator=g);rr_l=torch.randn(B,T,V,generator=g)
cl=torch.randint(0,V,(B,T),generator=g);rl=torch.randint(0,V,(B,T),generator=g)
cm=(torch.rand(B,T,generator=g)>0.3).float();rm=(torch.rand(B,T,generator=g)>0.3).float()
beta=0.1;ls=0.0
ref=R.dpo_sequence_loss(pc_l,pr_l,rc_l,rr_l,cl,rl,cm,rm,beta,ls)
print("ref",ref.item())
for norm in [False,True]:
pc=seqlp(pc_l,cl,cm,norm);pr=seqlp(pr_l,rl,rm,norm)
rc=seqlp(rc_l,cl,cm,norm);rr=seqlp(rr_l,rl,rm,norm)
logits=beta*((pc-rc)-(pr-rr))
loss=(-(1-ls)*torch.nn.functional.logsigmoid(logits)-ls*torch.nn.functional.logsigmoid(-logits)).mean()
print("norm",norm,loss.item())
EOF
python3 /app/_p21.pyref 0.7345190048217773 norm False 0.7345190048217773 norm True 0.7023628950119019 [stdout] ref 0.7345190048217773 norm False 0.7345190048217773 norm True 0.7023628950119019
Probe rloo_objective
cat > /app/_p22.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
g=torch.Generator().manual_seed(13)
B,T,V=6,3,5; gs=3
logits=torch.randn(B,T,V,generator=g);old=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g)
mask=(torch.rand(B,T,generator=g)>0.3).float()
rewards=torch.randn(B,generator=g)
cl,ch=0.2,0.3
ref=R.rloo_objective(logits,old,labels,mask,rewards,gs,cl,ch)
print("ref",ref.item())
# rloo advantages
adv=R.rloo_advantages(rewards,gs) # [B]
A=adv.unsqueeze(1).expand(B,T)
logp=tlp(logits,labels);oldlp=tlp(old,labels)
ratio=torch.exp(logp-oldlp)
surr1=ratio*A; surr2=torch.clamp(ratio,1-cl,1+ch)*A
pg=-torch.min(surr1,surr2)
print("plain mm", ((pg*mask).sum()/mask.sum()).item())
print("B*mm", B*((pg*mask).sum()/mask.sum()).item())
print("sum-batch of seqmean", ((pg*mask).sum(1)/mask.sum(1).clamp(min=1)).sum().item())
EOF
python3 /app/_p22.pyref 0.29768627882003784 plain mm 0.29768627882003784 B*mm 1.786117672920227 sum-batch of seqmean 1.788561224937439 [stdout] ref 0.29768627882003784 plain mm 0.29768627882003784 B*mm 1.786117672920227 sum-batch of seqmean 1.788561224937439
Probe grpo_objective
cat > /app/_p23.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
g=torch.Generator().manual_seed(17)
B,T,V=6,3,5; gs=3
logits=torch.randn(B,T,V,generator=g);old=torch.randn(B,T,V,generator=g);refl=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g)
mask=(torch.rand(B,T,generator=g)>0.3).float()
rewards=torch.randn(B,generator=g)
beta,cl,ch=0.05,0.2,0.2
def klpen(d,est):
if est=="k1": return d
if est=="k2": return 0.5*d*d
if est=="k3": return torch.exp(-d)+d-1
for est in ["k1","k2","k3"]:
ref=R.grpo_objective(logits,old,refl,labels,mask,rewards,gs,beta,cl,ch,True,est).item()
adv=R.grpo_advantages(rewards,gs,True).unsqueeze(1).expand(B,T)
logp=tlp(logits,labels);oldlp=tlp(old,labels);reflp=tlp(refl,labels)
ratio=torch.exp(logp-oldlp)
pg=-torch.min(ratio*adv,torch.clamp(ratio,1-cl,1+ch)*adv)
kl=klpen(logp-reflp,est)
per=pg+beta*kl
mm=(per*mask).sum()/mask.sum()
print(est,"ref",round(ref,5),"pg+bkl mm",round(mm.item(),5),"match",abs(ref-mm.item())<1e-5)
EOF
python3 /app/_p23.pyk1 ref 0.41223 pg+bkl mm 0.41223 match True k2 ref 0.46354 pg+bkl mm 0.46354 match True k3 ref 0.72606 pg+bkl mm 0.72606 match True [stdout] k1 ref 0.41223 pg+bkl mm 0.41223 match True k2 ref 0.46354 pg+bkl mm 0.46354 match True k3 ref 0.72606 pg+bkl mm 0.72606 match True
Probe ppo_objective
cat > /app/_p24.py << 'EOF'
import torch, reforacle as R
g=torch.Generator().manual_seed(19)
T=5
rewards=torch.randn(T,generator=g);values=torch.randn(T,generator=g);old_values=torch.randn(T,generator=g)
logp=torch.randn(T,generator=g)*0.1;old_logp=torch.randn(T,generator=g)*0.1
nv=torch.randn(1,generator=g)[0]
gamma,lam,cl,ch,vfc,vfcoef=0.99,0.95,0.2,0.2,0.2,0.5
ref=R.ppo_objective(rewards,values,old_values,logp,old_logp,nv,gamma,lam,cl,ch,vfc,vfcoef).item()
print("ref",ref)
gae=R.gae(rewards,values,nv,gamma,lam)
returns=gae+values
ratio=torch.exp(logp-old_logp)
def pgloss(adv):
return torch.max(-adv*ratio,-adv*torch.clamp(ratio,1-cl,1+ch)).mean()
vl=R.value_loss(values,old_values,returns,vfc)
for label,adv in [("raw",gae),("whiten",(gae-gae.mean())/(gae.std(unbiased=False)+1e-8)),
("whiten_ub",(gae-gae.mean())/(gae.std(unbiased=True)+1e-8))]:
tot=pgloss(adv)+vfcoef*vl
print(label,"pg",round(pgloss(adv).item(),5),"tot",round(tot.item(),5),"match",abs(tot.item()-ref)<1e-5)
print("vl",vl.item())
EOF
python3 /app/_p24.pyref 1.6760817766189575 raw pg 0.55052 tot 2.191 match False whiten pg 0.0356 tot 1.67608 match True whiten_ub pg 0.03184 tot 1.67232 match False vl 3.2809576988220215 [stdout] ref 1.6760817766189575 raw pg 0.55052 tot 2.191 match False whiten pg 0.0356 tot 1.67608 match True whiten_ub pg 0.03184 tot 1.67232 match False vl 3.2809576988220215
Final edge conventions
cat > /app/_p25.py << 'EOF'
import torch, reforacle as R
# normalize denom: std+eps vs sqrt(var+eps), big eps
x=torch.tensor([1.,2.,3.,4.])
print("norm eps1", R.normalize(x,1.0).tolist())
print(" (x-m)/(std_ub+1)", ((x-x.mean())/(x.std(unbiased=True)+1)).tolist())
print(" (x-m)/sqrt(var_ub+1)", ((x-x.mean())/torch.sqrt(x.var(unbiased=True)+1)).tolist())
# masked_max 2d
v=torch.tensor([[1.,5.,2.],[8.,3.,9.]]);m=torch.tensor([[1.,0.,1.],[0.,1.,1.]])
print("masked_max2d", R.masked_max(v,m))
# mode
print("mode 1122", R.mode_label(torch.tensor([1,1,2,2])), "torch.mode", torch.mode(torch.tensor([1,1,2,2])).values)
# smoothed_nll
logits=torch.tensor([[2.,1.,0.,-1.]]);labels=torch.tensor([0])
lp=logits.log_softmax(-1); nll=-lp[0,0]; K=4
s=0.1
print("smoothed ref", R.smoothed_nll(logits,labels,s).item())
print(" cand (1-s)*nll - s*mean(lp)", ((1-s)*nll - s*lp.mean()).item())
# top_p exact boundary
print("top_p 0.9", R.top_p_mask(torch.tensor([[0.5,0.3,0.15,0.05]]),0.9))
EOF
python3 /app/_p25.pynorm eps1 [-0.6547375321388245, -0.21824583411216736, 0.21824583411216736, 0.6547375321388245] (x-m)/(std_ub+1) [-0.6547375321388245, -0.21824583411216736, 0.21824583411216736, 0.6547375321388245] (x-m)/sqrt(var_ub+1) [-0.9185587167739868, -0.3061862289905548, 0.3061862289905548, 0.9185587167739868] masked_max2d tensor(9.) mode 1122 tensor(2) torch.mode tensor(1) smoothed ref 0.5901897549629211 cand (1-s)*nll - s*mean(lp) 0.5901897549629211 top_p 0.9 tensor([[0.5263, 0.3158, 0.1579, 0.0000]]) [stdout] norm eps1 [-0.6547375321388245, -0.21824583411216736, 0.21824583411216736, 0.6547375321388245] (x-m)/(std_ub+1) [-0.6547375321388245, -0.21824583411216736, 0.21824583411216736, 0.6547375321388245] (x-m)/sqrt(var_ub+1) [-0.9185587167739868, -0.3061862289905548, 0.3061862289905548, 0.9185587167739868] masked_max2d tensor(9.) mode 1122 tensor(2) torch.mode tensor(1) smoothed ref 0.5901897549629211 cand (1-s)*nll - s*mean(lp) 0.5901897549629211 top_p 0.9 tensor([[0.5263, 0.3158, 0.1579, 0.0000]])
from __future__ import annotations
import torch
# ---------------------------------------------------------------------------
# helpers
# ---------------------------------------------------------------------------
def _token_logprobs(logits, labels):
lp = torch.log_softmax(logits, dim=-1)
return lp.gather(-1, labels.unsqueeze(-1)).squeeze(-1)
def _masked_mean(values, mask, axis=None):
if axis is None:
num = (values * mask).sum()
den = mask.sum()
else:
num = (values * mask).sum(axis)
den = mask.sum(axis)
return num / den.clamp(min=1e-12)
# ---------------------------------------------------------------------------
# primitives
# ---------------------------------------------------------------------------
def masked_mean(values, mask, axis=None):
return _masked_mean(values, mask, axis)
def masked_sum(values, mask, axis=None):
if axis is None:
return (values * mask).sum()
return (values * mask).sum(axis)
def logsumexp(x, axis):
return torch.logsumexp(x, dim=axis)
def log_softmax(x, axis):
return torch.log_softmax(x, dim=axis)
def selective_logprobs(logits, labels, mask):
return (_token_logprobs(logits, labels) * mask).sum(-1)
def token_logprobs(logits, labels):
return _token_logprobs(logits, labels)
def sequence_logprob(logits, labels, mask, length_normalize):
lp = _token_logprobs(logits, labels)
s = (lp * mask).sum(-1)
if length_normalize:
return s / mask.sum(-1).clamp(min=1e-12)
return s
def entropy(logits, mask):
lp = torch.log_softmax(logits, dim=-1)
p = lp.exp()
ent = -(p * lp).sum(-1)
return _masked_mean(ent, mask)
def dpo_loss(pc, pr, rc, rr, beta, label_smoothing):
logits = beta * ((pc - rc) - (pr - rr))
loss = (-(1.0 - label_smoothing) * torch.nn.functional.logsigmoid(logits)
- label_smoothing * torch.nn.functional.logsigmoid(-logits))
return loss.mean()
def ipo_loss(pc, pr, rc, rr, beta):
h = (pc - rc) - (pr - rr)
loss = (h - 1.0 / (2.0 * beta)) ** 2
return loss.mean()
def grpo_advantages(rewards, group_size, scale_by_std):
g = rewards.view(-1, group_size)
mean = g.mean(dim=1, keepdim=True)
centered = g - mean
if scale_by_std:
std = g.std(dim=1, unbiased=False, keepdim=True)
centered = centered / (std + 1e-6)
return centered.reshape(rewards.shape)
def gae(rewards, values, next_value, gamma, lam):
T = rewards.shape[0]
adv = torch.zeros_like(rewards)
last = torch.zeros_like(next_value)
for t in range(T - 1, -1, -1):
v_next = next_value if t == T - 1 else values[t + 1]
delta = rewards[t] + gamma * v_next - values[t]
last = delta + gamma * lam * last
adv[t] = last
return adv
def kl_penalty(logp, ref_logp, estimator):
d = logp - ref_logp
if estimator == "k1":
return d
if estimator == "k2":
return 0.5 * d * d
# k3
return torch.exp(-d) + d - 1.0
def clipped_pg_loss(logp, old_logp, advantages, mask, clip_low, clip_high):
ratio = torch.exp(logp - old_logp)
unclipped = -advantages * ratio
clipped = -advantages * torch.clamp(ratio, 1.0 - clip_low, 1.0 + clip_high)
loss = torch.maximum(unclipped, clipped)
return _masked_mean(loss, mask)
def value_loss(values, old_values, returns, clip):
v_clipped = old_values + torch.clamp(values - old_values, -clip, clip)
l1 = (values - returns) ** 2
l2 = (v_clipped - returns) ** 2
return 0.5 * torch.maximum(l1, l2).mean()
def whiten(values, mask, shift_mean):
mean = _masked_mean(values, mask)
var = _masked_mean((values - mean) ** 2, mask)
whitened = (values - mean) / (torch.sqrt(var) + 1e-8)
if not shift_mean:
whitened = whitened + mean
return whitened
def dpo_sequence_loss(pc_logits, pr_logits, rc_logits, rr_logits,
chosen_labels, rejected_labels, chosen_mask, rejected_mask,
beta, label_smoothing):
pc = (_token_logprobs(pc_logits, chosen_labels) * chosen_mask).sum(-1)
pr = (_token_logprobs(pr_logits, rejected_labels) * rejected_mask).sum(-1)
rc = (_token_logprobs(rc_logits, chosen_labels) * chosen_mask).sum(-1)
rr = (_token_logprobs(rr_logits, rejected_labels) * rejected_mask).sum(-1)
return dpo_loss(pc, pr, rc, rr, beta, label_smoothing)
def grpo_objective(logits, old_logits, ref_logits, labels, completion_mask,
rewards, group_size, beta, clip_low, clip_high, scale_by_std,
kl_estimator):
adv = grpo_advantages(rewards, group_size, scale_by_std)
adv = adv.unsqueeze(-1).expand_as(labels.float() if labels.dtype.is_floating_point
else labels)
adv = adv.to(logits.dtype)
logp = _token_logprobs(logits, labels)
old_logp = _token_logprobs(old_logits, labels)
ref_logp = _token_logprobs(ref_logits, labels)
ratio = torch.exp(logp - old_logp)
unclipped = -adv * ratio
clipped = -adv * torch.clamp(ratio, 1.0 - clip_low, 1.0 + clip_high)
pg = torch.maximum(unclipped, clipped)
kl = kl_penalty(logp, ref_logp, kl_estimator)
return _masked_mean(pg + beta * kl, completion_mask)
def ppo_objective(rewards, values, old_values, logp, old_logp, next_value,
gamma, lam, clip_low, clip_high, vf_clip, vf_coef):
adv = gae(rewards, values, next_value, gamma, lam)
returns = adv + values
mean = adv.mean()
var = ((adv - mean) ** 2).mean()
adv_w = (adv - mean) / (torch.sqrt(var) + 1e-8)
ratio = torch.exp(logp - old_logp)
unclipped = -adv_w * ratio
clipped = -adv_w * torch.clamp(ratio, 1.0 - clip_low, 1.0 + clip_high)
pg = torch.maximum(unclipped, clipped).mean()
vl = value_loss(values, old_values, returns, vf_clip)
return pg + vf_coef * vl
def rloo_advantages(rewards, group_size):
g = rewards.view(-1, group_size)
total = g.sum(dim=1, keepdim=True)
baseline = (total - g) / (group_size - 1)
return (g - baseline).reshape(rewards.shape)
def reverse_kl(logp, ref_logp):
d = logp - ref_logp
return torch.exp(-d) + d - 1.0
def importance_ratio(logp, old_logp, clip):
ratio = torch.exp(logp - old_logp)
if clip is not None:
ratio = torch.clamp(ratio, 1.0 - clip, 1.0 + clip)
return ratio
def discounted_returns(rewards, gamma):
T = rewards.shape[0]
out = torch.zeros_like(rewards)
last = torch.zeros_like(rewards[0])
for t in range(T - 1, -1, -1):
last = rewards[t] + gamma * last
out[t] = last
return out
def normalize(x, eps):
mean = x.mean()
std = x.std(unbiased=True)
return (x - mean) / (std + eps)
def top_p_mask(probs, p):
sorted_probs, sorted_idx = torch.sort(probs, dim=-1, descending=True)
cumsum = sorted_probs.cumsum(dim=-1)
remove = cumsum >= p
# shift right so the token crossing the threshold is kept
remove[..., 1:] = remove[..., :-1].clone()
remove[..., 0] = False
sorted_probs = sorted_probs.masked_fill(remove, 0.0)
kept = torch.zeros_like(probs)
kept.scatter_(-1, sorted_idx, sorted_probs)
kept = kept / kept.sum(dim=-1, keepdim=True).clamp(min=1e-12)
return kept
def smoothed_nll(logits, labels, smoothing):
lp = torch.log_softmax(logits, dim=-1)
nll = -lp.gather(-1, labels.unsqueeze(-1)).squeeze(-1)
smooth = -lp.mean(dim=-1)
loss = (1.0 - smoothing) * nll + smoothing * smooth
return loss.mean()
def bradley_terry_logit(chosen_reward, rejected_reward, beta):
return beta * (chosen_reward - rejected_reward)
def rloo_objective(logits, old_logits, labels, mask, rewards, group_size, clip_low, clip_high):
adv = rloo_advantages(rewards, group_size)
adv = adv.unsqueeze(-1).to(logits.dtype)
logp = _token_logprobs(logits, labels)
old_logp = _token_logprobs(old_logits, labels)
ratio = torch.exp(logp - old_logp)
unclipped = -adv * ratio
clipped = -adv * torch.clamp(ratio, 1.0 - clip_low, 1.0 + clip_high)
pg = torch.maximum(unclipped, clipped)
return _masked_mean(pg, mask)
def reverse_kl_objective(logits, ref_logits, labels, mask, advantages, beta):
B = logits.shape[0]
logp = _token_logprobs(logits, labels)
ref_logp = _token_logprobs(ref_logits, labels)
d = logp - ref_logp
kl = torch.exp(-d) + d - 1.0
M = mask.sum().clamp(min=1e-12)
term_adv = -(((logp * mask).sum(0)) * advantages.sum(0)).sum() / M
term_kl = beta * B * ((kl * mask).sum() / M)
return term_adv + term_kl
def cross_entropy(logits, labels, ignore_index):
lp = torch.log_softmax(logits, dim=-1)
valid = labels != ignore_index
safe = labels.clamp(min=0)
nll = -lp.gather(-1, safe.unsqueeze(-1)).squeeze(-1)
nll = nll * valid
return nll.sum() / valid.sum().clamp(min=1e-12)
def top_k_mask(logits, k):
kth = torch.topk(logits, k, dim=-1).values[..., -1:]
return logits >= kth
def group_mean_baseline(rewards, group_size):
g = rewards.view(-1, group_size)
mean = g.mean(dim=1, keepdim=True)
return (g - mean).reshape(rewards.shape)
def lambda_returns(rewards, values, next_value, gamma, lam):
return gae(rewards, values, next_value, gamma, lam) + values
def symmetric_kl(logp, ref_logp):
d = logp - ref_logp
rev = torch.exp(-d) + d - 1.0
fwd = torch.exp(d) - d - 1.0
return 0.5 * (rev + fwd)
def huber_value_loss(values, returns, delta):
err = values - returns
a = err.abs()
loss = torch.where(a <= delta, 0.5 * err * err, delta * (a - 0.5 * delta))
return loss.mean()
def normalized_entropy(logits, mask):
V = logits.shape[-1]
return entropy(logits, mask) / torch.log(torch.tensor(float(V)))
def clip_fraction(logp, old_logp, clip):
ratio = torch.exp(logp - old_logp)
clipped = (ratio < 1.0 - clip) | (ratio > 1.0 + clip)
return clipped.float().mean()
def masked_whiten(values, mask, shift_mean):
mean = _masked_mean(values, mask)
var = _masked_mean((values - mean) ** 2, mask)
whitened = (values - mean) / (torch.sqrt(var) + 1e-8)
whitened = whitened * mask
if not shift_mean:
whitened = whitened + mean
return whitened
def logprob_at_temperature(logits, labels, temperature):
scaled = logits / temperature
return _token_logprobs(scaled, labels)
def advantage_mean_std(advantages, mask):
mean = _masked_mean(advantages, mask)
var = _masked_mean((advantages - mean) ** 2, mask)
std = torch.sqrt(var)
return torch.stack([mean, std])
def argmax_tokens(logits):
return logits.argmax(dim=-1)
def mode_label(labels):
vals, counts = torch.unique(labels, return_counts=True)
maxc = counts.max()
return vals[counts == maxc].max()
def median_reward(rewards):
return torch.quantile(rewards.float(), 0.5)
def masked_max(values, mask):
filled = values.masked_fill(mask == 0, float("-inf"))
return filled.max()
def quantile_lower(x, q):
return torch.quantile(x.float(), q, interpolation="lower")
def pad_mask_from_lengths(lengths, max_len):
ar = torch.arange(max_len, device=lengths.device)
return (ar.unsqueeze(0) < lengths.unsqueeze(1)).float()
def first_nonzero_index(mask):
nz = mask != 0
has = nz.any(dim=-1)
first = nz.float().argmax(dim=-1)
return torch.where(has, first, torch.full_like(first, -1))
def cumulative_max(x):
inclusive = torch.cummax(x, dim=0).values
return torch.cat([x[:1], inclusive[:-1]], dim=0)
def bucketize_reward(rewards, boundaries):
return torch.bucketize(rewards, boundaries, right=True)
Run full comparison harness
cat > /app/_harness.py << 'EOF'
import torch, reforacle as R
import importlib, objectives
importlib.reload(objectives)
O=objectives
g=torch.Generator().manual_seed(1234)
def rn(*s): return torch.randn(*s,generator=g)
def rint(hi,*s): return torch.randint(0,hi,s,generator=g)
def rmask(*s): return (torch.rand(*s,generator=g)>0.35).float()
results={}
def chk(name,args_fn,n=6,composite=False):
atol=1e-5 if composite else 1e-6; rtol=1e-4 if composite else 1e-5
ok=True; err=None
for _ in range(n):
args=args_fn()
try:
r=getattr(R,name)(*args); m=getattr(O,name)(*args)
except Exception as e:
ok=False; err=f"EXC {e}"; break
r=r if torch.is_tensor(r) else torch.tensor(r)
m=m if torch.is_tensor(m) else torch.tensor(m)
if r.shape!=m.shape: ok=False; err=f"shape {r.shape} vs {m.shape}"; break
if r.dtype!=m.dtype and not (r.is_floating_point() and m.is_floating_point()):
ok=False; err=f"dtype {r.dtype} vs {m.dtype}"; break
if not torch.allclose(r.float(),m.float(),atol=atol,rtol=rtol,equal_nan=True):
ok=False; err=f"val diff max {(r.float()-m.float()).abs().max().item():.2e}"; break
results[name]=(ok,err)
print(("PASS" if ok else "FAIL"),name, "" if ok else err)
BT=lambda:(torch.randint(2,5,(1,),generator=g).item(),torch.randint(2,5,(1,),generator=g).item())
def logits_labels_mask():
B,T=BT(); V=torch.randint(3,7,(1,),generator=g).item()
return rn(B,T,V),rint(V,B,T),rmask(B,T)
chk("masked_mean",lambda:(rn(4,5),rmask(4,5),None))
chk("masked_mean",lambda:(rn(4,5),rmask(4,5),1))
chk("masked_sum",lambda:(rn(4,5),rmask(4,5),0))
chk("logsumexp",lambda:(rn(3,5),1))
chk("log_softmax",lambda:(rn(3,5),-1))
chk("selective_logprobs",logits_labels_mask)
chk("token_logprobs",lambda:(lambda x:(x[0],x[1]))(logits_labels_mask()))
chk("sequence_logprob",lambda:(*logits_labels_mask(),bool(torch.randint(0,2,(1,),generator=g).item())))
chk("entropy",logits_labels_mask)
def dpo_args():
B=3; return rn(B),rn(B),rn(B),rn(B),torch.rand(1,generator=g).item()+0.05,torch.rand(1,generator=g).item()*0.2
chk("dpo_loss",dpo_args,composite=True)
chk("ipo_loss",lambda:(rn(3),rn(3),rn(3),rn(3),torch.rand(1,generator=g).item()+0.05),composite=True)
chk("grpo_advantages",lambda:(rn(12),3,bool(torch.randint(0,2,(1,),generator=g).item())))
chk("gae",lambda:(rn(6),rn(6),rn(1)[0],0.99,0.95))
chk("kl_penalty",lambda:(rn(8),rn(8),["k1","k2","k3"][torch.randint(0,3,(1,),generator=g).item()]))
chk("clipped_pg_loss",lambda:(rn(4,5),rn(4,5),rn(4,5),rmask(4,5),0.2,0.3),composite=True)
chk("value_loss",lambda:(rn(6),rn(6),rn(6),0.2),composite=True)
chk("whiten",lambda:(rn(7),rmask(7),bool(torch.randint(0,2,(1,),generator=g).item())))
def dpo_seq_args():
B,T=3,4; V=5
return (rn(B,T,V),rn(B,T,V),rn(B,T,V),rn(B,T,V),rint(V,B,T),rint(V,B,T),rmask(B,T),rmask(B,T),0.1,0.05)
chk("dpo_sequence_loss",dpo_seq_args,composite=True)
def grpo_obj_args():
B,T,V=6,3,5
return (rn(B,T,V),rn(B,T,V),rn(B,T,V),rint(V,B,T),rmask(B,T),rn(B),3,0.05,0.2,0.2,
bool(torch.randint(0,2,(1,),generator=g).item()),["k1","k2","k3"][torch.randint(0,3,(1,),generator=g).item()])
chk("grpo_objective",grpo_obj_args,composite=True)
chk("ppo_objective",lambda:(rn(6),rn(6),rn(6),rn(6)*0.1,rn(6)*0.1,rn(1)[0],0.99,0.95,0.2,0.2,0.2,0.5),composite=True)
chk("rloo_advantages",lambda:(rn(12),4))
chk("reverse_kl",lambda:(rn(8),rn(8)))
chk("importance_ratio",lambda:(rn(6)*0.2,rn(6)*0.2,[None,0.2][torch.randint(0,2,(1,),generator=g).item()]))
chk("discounted_returns",lambda:(rn(7),0.9))
chk("normalize",lambda:(rn(8),1e-8))
chk("top_p_mask",lambda:(torch.softmax(rn(4,6),-1),0.7))
chk("smoothed_nll",lambda:(rn(5,7),rint(7,5),0.1),composite=True)
chk("bradley_terry_logit",lambda:(rn(4),rn(4),0.5))
def rloo_obj_args():
B,T,V=6,3,5
return (rn(B,T,V),rn(B,T,V),rint(V,B,T),rmask(B,T),rn(B),3,0.2,0.3)
chk("rloo_objective",rloo_obj_args,composite=True)
def rkl_obj_args():
B,T,V=4,3,5
return (rn(B,T,V),rn(B,T,V),rint(V,B,T),rmask(B,T),rn(B,T),0.1)
chk("reverse_kl_objective",rkl_obj_args,composite=True)
chk("cross_entropy",lambda:(rn(5,7),rint(7,5),-100))
chk("top_k_mask",lambda:(rn(4,7),2))
chk("group_mean_baseline",lambda:(rn(12),4))
chk("lambda_returns",lambda:(rn(6),rn(6),rn(1)[0],0.99,0.95))
chk("symmetric_kl",lambda:(rn(8),rn(8)))
chk("huber_value_loss",lambda:(rn(6),rn(6),1.0),composite=True)
chk("normalized_entropy",logits_labels_mask)
chk("clip_fraction",lambda:(rn(8)*0.3,rn(8)*0.3,0.2))
chk("masked_whiten",lambda:(rn(7),rmask(7),bool(torch.randint(0,2,(1,),generator=g).item())))
chk("logprob_at_temperature",lambda:(lambda x:(x[0],x[1],1.5))(logits_labels_mask()))
chk("advantage_mean_std",lambda:(rn(4,5),rmask(4,5)))
chk("argmax_tokens",lambda:(rn(3,4,6),))
chk("mode_label",lambda:(rint(4,10),))
chk("median_reward",lambda:(rn(7),))
chk("masked_max",lambda:(rn(4,5),rmask(4,5)))
chk("quantile_lower",lambda:(rn(9),0.3))
chk("pad_mask_from_lengths",lambda:(torch.randint(1,5,(4,),generator=g),5))
chk("first_nonzero_index",lambda:(rmask(4,6),))
chk("cumulative_max",lambda:(rn(6),))
chk("bucketize_reward",lambda:(rn(8),torch.tensor([-1.,0.,1.])))
print("\n== SUMMARY ==")
fails=[k for k,(ok,_) in results.items() if not ok]
print(f"{sum(1 for _,(ok,_) in results.items() if ok)}/{len(results)} passing")
if fails: print("FAILS:",fails)
EOF
python3 /app/_harness.pyPASS masked_mean PASS masked_mean PASS masked_sum PASS logsumexp PASS log_softmax PASS selective_logprobs PASS token_logprobs PASS sequence_logprob FAIL entropy EXC entropy() takes exactly 2 positional arguments (3 given) PASS dpo_loss PASS ipo_loss PASS grpo_advantages PASS gae PASS kl_penalty PASS clipped_pg_loss PASS value_loss FAIL whiten val diff max 2.07e+08 PASS dpo_sequence_loss PASS grpo_objective PASS ppo_objective PASS rloo_advantages PASS reverse_kl PASS importance_ratio PASS discounted_returns PASS normalize PASS top_p_mask PASS smoothed_nll PASS bradley_terry_logit PASS rloo_objective PASS reverse_kl_objective PASS cross_entropy PASS top_k_mask PASS group_mean_baseline PASS lambda_returns PASS symmetric_kl PASS huber_value_loss FAIL normalized_entropy EXC normalized_entropy() takes exactly 2 positional arguments (3 given) PASS clip_fraction PASS masked_whiten PASS logprob_at_temperature PASS advantage_mean_std PASS argmax_tokens PASS mode_label PASS median_reward PASS masked_max PASS quantile_lower PASS pad_mask_from_lengths PASS first_nonzero_index PASS cumulative_max PASS bucketize_reward == SUMMARY == 46/49 passing FAILS: ['entropy', 'whiten', 'normalized_entropy'] [stdout] PASS masked_mean PASS masked_mean PASS masked_sum PASS logsumexp PASS log_softmax PASS selective_logprobs PASS token_logprobs PASS sequence_logprob FAIL entropy EXC entropy() takes exactly 2 positional arguments (3 given) PASS dpo_loss PASS ipo_loss PASS grpo_advantages PASS gae PASS kl_penalty PASS clipped_pg_loss PASS value_loss FAIL whiten val diff max 2.07e+08 PASS dpo_sequence_loss PASS grpo_objective PASS ppo_objective PASS rloo_advantages PASS reverse_kl PASS importance_ratio PASS discounted_returns PASS normalize PASS top_p_mask PASS smoothed_nll PASS bradley_terry_logit PASS rloo_objective PASS reverse_kl_objective PASS cross_entropy PASS top_k_mask PASS group_mean_baseline PASS lambda_returns PASS symmetric_kl PASS huber_value_loss FAIL normalized_entropy EXC normalized_entropy() takes exactly 2 positional arguments (3 given) PASS clip_fraction PASS masked_whiten PASS logprob_at_temperature PASS advantage_mean_std PASS argmax_tokens PASS mode_label PASS median_reward PASS masked_max PASS quantile_lower PASS pad_mask_from_lengths PASS first_nonzero_index PASS cumulative_max PASS bucketize_reward == SUMMARY == 46/49 passing FAILS: ['entropy', 'whiten', 'normalized_entropy']
Find whiten failing case (empty mask)
cat > /app/_dbg.py << 'EOF'
import torch, reforacle as R, objectives as O
g=torch.Generator().manual_seed(1234)
def rn(*s): return torch.randn(*s,generator=g)
def rmask(*s): return (torch.rand(*s,generator=g)>0.35).float()
# replicate harness sequence up to whiten... instead just search for failing case
for i in range(200):
v=rn(7); m=rmask(7); sh=bool(torch.randint(0,2,(1,),generator=g).item())
if m.sum()==0:
r=R.whiten(v,m,sh); print("empty mask r=",r);
print("mine", O.whiten(v,m,sh)); break
EOF
python3 /app/_dbg.py(Bash completed with no output)
Reproduce whiten failure
cat > /app/_dbg.py << 'EOF'
import torch, reforacle as R, objectives as O
g=torch.Generator().manual_seed(55)
def rn(*s): return torch.randn(*s,generator=g)
def rmask(*s): return (torch.rand(*s,generator=g)>0.35).float()
for i in range(500):
v=rn(7); m=rmask(7); sh=bool(torch.randint(0,2,(1,),generator=g).item())
r=R.whiten(v,m,sh); mine=O.whiten(v,m,sh)
if not torch.allclose(r,mine,atol=1e-5,rtol=1e-4,equal_nan=True):
print("FAIL sh",sh,"masksum",m.sum().item())
print("v",v.tolist())
print("m",m.tolist())
print("ref ",r.tolist())
print("mine",mine.tolist())
break
else: print("no fail in 500")
EOF
python3 /app/_dbg.pyFAIL sh False masksum 1.0 v [-0.8841636180877686, 0.5962507724761963, -0.7032215595245361, -0.6375294327735901, -0.8732007741928101, 1.6124345064163208, -2.046596050262451] m [0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0] ref [-2466.979248046875, 12337.1640625, -657.5587768554688, -0.6375294327735901, -2357.350830078125, 22499.00390625, -14091.3046875] mine [-24663418.0, 123378016.0, -6569213.0, -0.6375294327735901, -23567134.0, 224996400.0, -140906672.0] [stdout] FAIL sh False masksum 1.0 v [-0.8841636180877686, 0.5962507724761963, -0.7032215595245361, -0.6375294327735901, -0.8732007741928101, 1.6124345064163208, -2.046596050262451] m [0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0] ref [-2466.979248046875, 12337.1640625, -657.5587768554688, -0.6375294327735901, -2357.350830078125, 22499.00390625, -14091.3046875] mine [-24663418.0, 123378016.0, -6569213.0, -0.6375294327735901, -23567134.0, 224996400.0, -140906672.0]
/app/objectives.py
mean = _masked_mean(values, mask)
var = _masked_mean((values - mean) ** 2, mask)
whitened = (values - mean) / (torch.sqrt(var) + 1e-8)
if not shift_mean:
whitened = whitened + mean
return whitened mean = _masked_mean(values, mask)
var = _masked_mean((values - mean) ** 2, mask)
whitened = (values - mean) / torch.sqrt(var + 1e-8)
if not shift_mean:
whitened = whitened + mean
return whitened/app/objectives.py
mean = _masked_mean(values, mask)
var = _masked_mean((values - mean) ** 2, mask)
whitened = (values - mean) / (torch.sqrt(var) + 1e-8)
whitened = whitened * mask mean = _masked_mean(values, mask)
var = _masked_mean((values - mean) ** 2, mask)
whitened = (values - mean) / torch.sqrt(var + 1e-8)
whitened = whitened * mask/app/objectives.py
adv_w = (adv - mean) / (torch.sqrt(var) + 1e-8)
adv_w = (adv - mean) / torch.sqrt(var + 1e-8)
Rerun harness after fixes
sed -i 's/^chk("entropy",logits_labels_mask)/chk("entropy",lambda:(lambda x:(x[0],x[2]))(logits_labels_mask()))/' _harness.py
sed -i 's/^chk("normalized_entropy",logits_labels_mask)/chk("normalized_entropy",lambda:(lambda x:(x[0],x[2]))(logits_labels_mask()))/' _harness.py
python3 _harness.py 2>&1 | tail -6PASS first_nonzero_index PASS cumulative_max PASS bucketize_reward == SUMMARY == 49/49 passing [stdout] PASS first_nonzero_index PASS cumulative_max PASS bucketize_reward == SUMMARY == 49/49 passing
Edge-case stress test
cat > /app/_edge.py << 'EOF'
import torch, reforacle as R, objectives as O
fail=[]
def eq(name,r,m,composite=False):
atol=1e-5 if composite else 1e-6; rtol=1e-4 if composite else 1e-5
r=r if torch.is_tensor(r) else torch.tensor(r); m=m if torch.is_tensor(m) else torch.tensor(m)
if r.shape!=m.shape or not torch.allclose(r.float(),m.float(),atol=atol,rtol=rtol,equal_nan=True):
fail.append((name,r,m))
print("FAIL",name,"\n ref",r,"\n mine",m)
def C(name,*a,comp=False):
try: eq(name,getattr(R,name)(*a),getattr(O,name)(*a),comp)
except Exception as e: fail.append((name,e)); print("EXC",name,e)
# all-zero mask
z=torch.zeros(5); C("masked_mean",torch.randn(5),z,None); C("masked_sum",torch.randn(5),z,None)
C("masked_max",torch.randn(5),z) # all masked -> -inf
C("whiten",torch.randn(5),z,True); C("masked_whiten",torch.randn(5),z,False)
C("advantage_mean_std",torch.randn(5),z)
# all-ones mask
o=torch.ones(4,3); C("masked_mean",torch.randn(4,3),o,1)
# constant input normalize/whiten
C("normalize",torch.full((5,),3.0),1e-8)
C("whiten",torch.full((5,),2.0),torch.ones(5),True)
# grpo/rloo group sizes
r=torch.randn(20)
for gs in [1,2,4,5,10,20]:
C("grpo_advantages",r,gs,True); C("grpo_advantages",r,gs,False)
C("rloo_advantages",r,gs); C("group_mean_baseline",r,gs)
# top_p boundaries
p=torch.softmax(torch.randn(3,6),-1)
for pp in [0.0,0.01,0.5,0.99,1.0]: C("top_p_mask",p,pp)
# top_k
lg=torch.randn(4,7)
for k in [1,2,7]: C("top_k_mask",lg,k)
C("top_k_mask",torch.tensor([[1.,1.,1.,1.]]),2)
# quantile edges
x=torch.randn(9)
for q in [0.0,0.25,0.5,0.75,1.0]: C("quantile_lower",x,q)
C("median_reward",torch.randn(6)); C("median_reward",torch.randn(7))
# cross_entropy with ignores
lg=torch.randn(6,5); lb=torch.randint(0,5,(6,)); lb[::2]=-100
C("cross_entropy",lg,lb,-100)
C("cross_entropy",torch.randn(3,4),torch.tensor([-100,-100,-100]),-100)
C("smoothed_nll",torch.randn(5,6),torch.randint(0,6,(5,)),0.0,comp=True)
C("smoothed_nll",torch.randn(5,6),torch.randint(0,6,(5,)),0.3,comp=True)
# bucketize exact boundaries + right
C("bucketize_reward",torch.tensor([-1.,0.,0.5,1.,2.,3.,5.]),torch.tensor([0.,1.,2.,3.]))
# first_nonzero all zero / all nonzero
C("first_nonzero_index",torch.tensor([[0.,0.,0.],[1.,1.,1.],[0.,0.,2.]]))
# cumulative_max 2d and negatives
C("cumulative_max",torch.tensor([[-1.,5.],[3.,-2.],[0.,9.]]))
C("cumulative_max",torch.randn(8))
# mode ties/unique
C("mode_label",torch.tensor([3,3,1,1,2])); C("mode_label",torch.tensor([5])); C("mode_label",torch.tensor([1,2,3,4]))
# pad mask
C("pad_mask_from_lengths",torch.tensor([0,1,5,3]),5)
# importance_ratio None + clip
C("importance_ratio",torch.randn(6),torch.randn(6),None)
C("importance_ratio",torch.randn(6),torch.randn(6),0.1)
# kl estimators
for est in ["k1","k2","k3"]: C("kl_penalty",torch.randn(7),torch.randn(7),est)
# sequence_logprob length norm with zero-length seq
lg=torch.randn(2,3,4); lb=torch.randint(0,4,(2,3)); mk=torch.tensor([[0.,0.,0.],[1.,1.,0.]])
C("sequence_logprob",lg,lb,mk,True); C("sequence_logprob",lg,lb,mk,False)
C("entropy",lg,mk); C("normalized_entropy",lg,mk); C("selective_logprobs",lg,lb,mk)
# gae/discounted 2d
C("discounted_returns",torch.randn(4,3),0.9)
C("value_loss",torch.randn(6),torch.randn(6),torch.randn(6),0.0,comp=True)
C("huber_value_loss",torch.randn(6),torch.randn(6),0.5,comp=True)
C("clip_fraction",torch.randn(10)*0.5,torch.randn(10)*0.5,0.2)
print("\nDONE fails:",len(fail))
EOF
python3 /app/_edge.pyFAIL masked_max ref tensor(0.) mine tensor(-inf) DONE fails: 1 [stdout] FAIL masked_max ref tensor(0.) mine tensor(-inf) DONE fails: 1
/app/objectives.py
filled = values.masked_fill(mask == 0, float("-inf"))
return filled.max() filled = values.masked_fill(mask == 0, float("-inf"))
res = filled.max()
return torch.where(torch.isinf(res), torch.zeros_like(res), res)Verify fix and multi-seed harness
python3 _edge.py 2>&1 | tail -3 && echo "--- multi-seed main harness ---" && for s in 7 42 99 2024; do sed "s/manual_seed(1234)/manual_seed($s)/" _harness.py > _h_$s.py; python3 _h_$s.py 2>&1 | tail -1; rm _h_$s.py; done
DONE fails: 0 --- multi-seed main harness --- 49/49 passing 49/49 passing 49/49 passing 49/49 passing [stdout] DONE fails: 0 --- multi-seed main harness --- 49/49 passing 49/49 passing 49/49 passing 49/49 passing
/app/objectives.py
adv = grpo_advantages(rewards, group_size, scale_by_std)
adv = adv.unsqueeze(-1).expand_as(labels.float() if labels.dtype.is_floating_point
else labels)
adv = adv.to(logits.dtype)
logp = _token_logprobs(logits, labels) adv = grpo_advantages(rewards, group_size, scale_by_std).unsqueeze(-1)
logp = _token_logprobs(logits, labels)Final verification and standalone check
python3 _harness.py 2>&1 | tail -2 && python3 _edge.py 2>&1 | tail -1
# confirm standalone (no reforacle dependency) by importing in isolation
python3 -c "import ast; ast.parse(open('/app/objectives.py').read()); print('syntax OK')"
grep -c "reforacle" /app/objectives.py && echo "references present" || echo "no reforacle references (good)"== SUMMARY == 49/49 passing DONE fails: 0 syntax OK 0 no reforacle references (good) [stdout] == SUMMARY == 49/49 passing DONE fails: 0 syntax OK 0 no reforacle references (good)
Remove temporary probe files
rm -f /app/_p*.py /app/_dbg.py /app/_harness.py /app/_edge.py; ls /app
__pycache__ objectives.py oracle probe.py reforacle.cpython-312-aarch64-linux-gnu.so timer.sh [stdout] __pycache__ objectives.py oracle probe.py reforacle.cpython-312-aarch64-linux-gnu.so timer.sh
No file changes, the agent left the workspace untouched.
=== SyncVals verifier (post-training objectives, sealed oracle) === FUNCTIONS PASSED: 42/49 (partial_score=0.8571) REWARD = 0 (1 iff ALL 49 functions match the sealed reference exactly) The agent probes a sealed compiled oracle (reforacle) to recover 49 post-training objective conventions (DPO/IPO/GRPO/PPO/GAE/KL/RLOO) and reimplements them in torch.
Reproduce this trial: git checkout 2f94510 && PYTHONPATH=src python3 scripts/build_site.py , then open trial/trial_67a684402caa49b1. Re-running the agent live requires EVAL_PLATFORM_ENABLE_OAUTH_SMOKE=1 and is non-deterministic.
Trial trial_67a684402caa49b1 · verifier authoritative; classifier explanatory.