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:59List reforacle functions
python3 -c "import reforacle; print([x for x in dir(reforacle) if not x.startswith('_')])"['F', 'advantage_mean_std', 'argmax_tokens', 'bradley_terry_logit', 'bucketize_reward', 'clip_fraction', 'clipped_pg_loss', 'cross_entropy', 'cumulative_max', 'discounted_returns', 'dpo_loss', 'dpo_sequence_loss', 'entropy', 'first_nonzero_index', 'gae', 'group_mean_baseline', 'grpo_advantages', 'grpo_objective', 'huber_value_loss', 'importance_ratio', 'ipo_loss', 'kl_penalty', 'lambda_returns', 'log_softmax', 'logprob_at_temperature', 'logsumexp', 'masked_max', 'masked_mean', 'masked_sum', 'masked_whiten', 'median_reward', 'mode_label', 'normalize', 'normalized_entropy', 'pad_mask_from_lengths', 'ppo_objective', 'quantile_lower', 'reverse_kl', 'reverse_kl_objective', 'rloo_advantages', 'rloo_objective', 'selective_logprobs', 'sequence_logprob', 'smoothed_nll', 'symmetric_kl', 'token_logprobs', 'top_k_mask', 'top_p_mask', 'torch', 'value_loss', 'whiten'] [stdout] ['F', 'advantage_mean_std', 'argmax_tokens', 'bradley_terry_logit', 'bucketize_reward', 'clip_fraction', 'clipped_pg_loss', 'cross_entropy', 'cumulative_max', 'discounted_returns', 'dpo_loss', 'dpo_sequence_loss', 'entropy', 'first_nonzero_index', 'gae', 'group_mean_baseline', 'grpo_advantages', 'grpo_objective', 'huber_value_loss', 'importance_ratio', 'ipo_loss', 'kl_penalty', 'lambda_returns', 'log_softmax', 'logprob_at_temperature', 'logsumexp', 'masked_max', 'masked_mean', 'masked_sum', 'masked_whiten', 'median_reward', 'mode_label', 'normalize', 'normalized_entropy', 'pad_mask_from_lengths', 'ppo_objective', 'quantile_lower', 'reverse_kl', 'reverse_kl_objective', 'rloo_advantages', 'rloo_objective', 'selective_logprobs', 'sequence_logprob', 'smoothed_nll', 'symmetric_kl', 'token_logprobs', 'top_k_mask', 'top_p_mask', 'torch', 'value_loss', 'whiten']
Probe masked functions
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(1)
v=torch.randn(6,generator=g)
m=(torch.rand(6,generator=g)>0.4).float()
print('v',v.tolist())
print('m',m.tolist())
print('masked_mean none',R.masked_mean(v,m).item())
print('masked_sum none',R.masked_sum(v,m).item())
# 2d
V=torch.randn(3,4,generator=g); M=(torch.rand(3,4,generator=g)>0.4).float()
print('mm axis1',R.masked_mean(V,M,axis=1).tolist())
print('ms axis1',R.masked_sum(V,M,axis=1).tolist())
print('mm axis0',R.masked_mean(V,M,axis=0).tolist())
# empty row mean
M2=M.clone(); M2[0]=0
print('mm empty row',R.masked_mean(V,M2,axis=1).tolist())
"
v [0.6613521575927734, 0.266924113035202, 0.06167725846171379, 0.6213173270225525, -0.4519059658050537, -0.16613022983074188] m [1.0, 0.0, 1.0, 1.0, 1.0, 1.0] masked_mean none 0.14526210725307465 masked_sum none 0.7263105511665344 mm axis1 [0.07571375370025635, 0.6249607801437378, -0.5403615832328796] ms axis1 [0.22714126110076904, 1.2499215602874756, -0.5403615832328796] mm axis0 [-1.194824457168579, 0.15644630789756775, -0.5403615832328796, 1.1794973611831665] mm empty row [0.0, 0.6249607801437378, -0.5403615832328796] [stdout] v [0.6613521575927734, 0.266924113035202, 0.06167725846171379, 0.6213173270225525, -0.4519059658050537, -0.16613022983074188] m [1.0, 0.0, 1.0, 1.0, 1.0, 1.0] masked_mean none 0.14526210725307465 masked_sum none 0.7263105511665344 mm axis1 [0.07571375370025635, 0.6249607801437378, -0.5403615832328796] ms axis1 [0.22714126110076904, 1.2499215602874756, -0.5403615832328796] mm axis0 [-1.194824457168579, 0.15644630789756775, -0.5403615832328796, 1.1794973611831665] mm empty row [0.0, 0.6249607801437378, -0.5403615832328796]
Probe masked eps
python3 -c "
import torch, reforacle as R
v=torch.tensor([1.0,1.0,1.0,5.0])
m=torch.tensor([1.0,1.0,1.0,0.0])
print('mean ones',repr(R.masked_mean(v,m).item())) # 1.0?
print('sum',repr(R.masked_sum(v,m).item()))
# axis default None on 2d -> scalar over all?
V=torch.tensor([[1.0,2.0],[3.0,4.0]]); M=torch.tensor([[1.0,0.0],[1.0,1.0]])
print('mean None 2d',repr(R.masked_mean(V,M).item()))
print('sum None 2d',repr(R.masked_sum(V,M).item()))
"
mean ones 1.0 sum 3.0 mean None 2d 2.6666667461395264 sum None 2d 8.0 [stdout] mean ones 1.0 sum 3.0 mean None 2d 2.6666667461395264 sum None 2d 8.0
Probe logsumexp/softmax/token_logprobs
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(2)
x=torch.randn(3,5,generator=g)
print('logsumexp ax1',R.logsumexp(x,1).tolist())
print('torch lse',torch.logsumexp(x,1).tolist())
print('log_softmax ax1 close', torch.allclose(R.log_softmax(x,1), torch.log_softmax(x,1)))
# token_logprobs
logits=torch.randn(2,4,6,generator=g)
labels=torch.randint(0,6,(2,4),generator=g)
tl=R.token_logprobs(logits,labels)
print('token_logprobs shape',tl.shape, tl.tolist())
lsm=torch.log_softmax(logits,-1)
print('match', torch.allclose(tl, lsm.gather(-1,labels.unsqueeze(-1)).squeeze(-1)))
"
logsumexp ax1 [1.8161442279815674, 1.6194981336593628, 1.9888298511505127] torch lse [1.8161442279815674, 1.6194981336593628, 1.9888298511505127] log_softmax ax1 close True token_logprobs shape torch.Size([2, 4]) [[-3.0128235816955566, -3.100942850112915, -1.4437549114227295, -3.372281074523926], [-2.2239389419555664, -3.2668118476867676, -3.0958871841430664, -2.9098703861236572]] match True [stdout] logsumexp ax1 [1.8161442279815674, 1.6194981336593628, 1.9888298511505127] torch lse [1.8161442279815674, 1.6194981336593628, 1.9888298511505127] log_softmax ax1 close True token_logprobs shape torch.Size([2, 4]) [[-3.0128235816955566, -3.100942850112915, -1.4437549114227295, -3.372281074523926], [-2.2239389419555664, -3.2668118476867676, -3.0958871841430664, -2.9098703861236572]] match True
Probe selective_logprobs, sequence_logprob
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(3)
logits=torch.randn(2,4,6,generator=g)
labels=torch.randint(0,6,(2,4),generator=g)
mask=(torch.rand(2,4,generator=g)>0.4).float()
sl=R.selective_logprobs(logits,labels,mask)
print('selective_logprobs shape',sl.shape)
print(sl.tolist())
lsm=torch.log_softmax(logits,-1).gather(-1,labels.unsqueeze(-1)).squeeze(-1)
print('lp*mask', (lsm*mask).tolist())
# sequence_logprob
for ln in [False,True]:
print('seq ln',ln, R.sequence_logprob(logits,labels,mask,ln).tolist())
print('sum lp*mask',(lsm*mask).sum(-1).tolist())
print('mask sum',mask.sum(-1).tolist())
"
selective_logprobs shape torch.Size([2]) [-3.958949565887451, -1.1747310161590576] lp*mask [[-0.0, -2.5062005519866943, -1.4527490139007568, -0.0], [-0.8520113825798035, -0.3227196931838989, -0.0, -0.0]] seq ln False [-3.958949565887451, -1.1747310161590576] seq ln True [-1.9794747829437256, -0.5873655080795288] sum lp*mask [-3.958949565887451, -1.1747310161590576] mask sum [2.0, 2.0] [stdout] selective_logprobs shape torch.Size([2]) [-3.958949565887451, -1.1747310161590576] lp*mask [[-0.0, -2.5062005519866943, -1.4527490139007568, -0.0], [-0.8520113825798035, -0.3227196931838989, -0.0, -0.0]] seq ln False [-3.958949565887451, -1.1747310161590576] seq ln True [-1.9794747829437256, -0.5873655080795288] sum lp*mask [-3.958949565887451, -1.1747310161590576] mask sum [2.0, 2.0]
Probe entropy
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(4)
logits=torch.randn(2,4,6,generator=g)
mask=(torch.rand(2,4,generator=g)>0.4).float()
e=R.entropy(logits,mask)
print('entropy',e.shape, e.tolist() if e.dim() else e.item())
# per-token entropy
p=torch.softmax(logits,-1); lsm=torch.log_softmax(logits,-1)
ent=-(p*lsm).sum(-1)
print('ent pertoken',ent.tolist())
print('masked mean of ent', (ent*mask).sum()/mask.sum())
print('per seq masked mean',((ent*mask).sum(-1)/mask.sum(-1)).tolist())
ne=R.normalized_entropy(logits,mask)
print('normalized_entropy',ne.shape, ne)
print('log6',torch.log(torch.tensor(6.0)))
"
entropy torch.Size([]) 1.5284241437911987 ent pertoken [[1.5357071161270142, 1.5358145236968994, 1.17733895778656, 1.6326051950454712], [1.588403344154358, 1.1268160343170166, 1.286355972290039, 1.3568737506866455]] masked mean of ent tensor(1.5284) per seq masked mean [1.58420991897583, 1.4726386070251465] normalized_entropy torch.Size([]) tensor(0.8530) log6 tensor(1.7918) [stdout] entropy torch.Size([]) 1.5284241437911987 ent pertoken [[1.5357071161270142, 1.5358145236968994, 1.17733895778656, 1.6326051950454712], [1.588403344154358, 1.1268160343170166, 1.286355972290039, 1.3568737506866455]] masked mean of ent tensor(1.5284) per seq masked mean [1.58420991897583, 1.4726386070251465] normalized_entropy torch.Size([]) tensor(0.8530) log6 tensor(1.7918)
Probe dpo/ipo
python3 -c "
import torch, reforacle as R
import torch.nn.functional as F
g=torch.Generator().manual_seed(5)
pc=torch.randn(4,generator=g); pr=torch.randn(4,generator=g)
rc=torch.randn(4,generator=g); rr=torch.randn(4,generator=g)
beta=0.1
for ls in [0.0,0.1]:
print('dpo ls',ls, R.dpo_loss(pc,pr,rc,rr,beta,ls).item())
logits=beta*((pc-rc)-(pr-rr))
# standard dpo: -logsigmoid(beta*((pc-pr)-(rc-rr)))
li=beta*((pc-pr)-(rc-rr))
print('std dpo mean', (-F.logsigmoid(li)).mean().item())
print('std dpo ls0.1', ((-F.logsigmoid(li))*0.9 + (-F.logsigmoid(-li))*0.1).mean().item())
print('ipo', R.ipo_loss(pc,pr,rc,rr,beta).item())
print('ipo formula', ((li - 1/(2*beta))**2).mean().item())
"
dpo ls 0.0 0.6920198202133179 dpo ls 0.1 0.6929887533187866 std dpo mean 0.6920198202133179 std dpo ls0.1 0.6929887533187866 ipo 27.01028823852539 ipo formula 24.93289566040039 [stdout] dpo ls 0.0 0.6920198202133179 dpo ls 0.1 0.6929887533187866 std dpo mean 0.6920198202133179 std dpo ls0.1 0.6929887533187866 ipo 27.01028823852539 ipo formula 24.93289566040039
Probe ipo formula
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(5)
pc=torch.randn(4,generator=g); pr=torch.randn(4,generator=g)
rc=torch.randn(4,generator=g); rr=torch.randn(4,generator=g)
beta=0.1
h=(pc-pr)-(rc-rr)
print('ipo', R.ipo_loss(pc,pr,rc,rr,beta).item())
print('h - 1/2b', ((h - 1/(2*beta))**2).mean().item())
print('beta*h', ((beta*h - 1/(2*beta))**2).mean().item())
"
ipo 27.01028823852539 h - 1/2b 27.01028823852539 beta*h 24.93289566040039 [stdout] ipo 27.01028823852539 h - 1/2b 27.01028823852539 beta*h 24.93289566040039
Probe advantages
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(6)
rewards=torch.randn(6,generator=g)
print('rewards',rewards.tolist())
for gs in [3]:
for sc in [False,True]:
print('grpo gs',gs,'scale',sc, R.grpo_advantages(rewards,gs,sc).tolist())
print('rloo', R.rloo_advantages(rewards,gs).tolist())
print('gmb', R.group_mean_baseline(rewards,gs).tolist())
"
rewards [-1.8743985891342163, -0.9936632513999939, 0.7184672355651855, -0.6984538435935974, -1.471555233001709, 0.1777002215385437] grpo gs 3 scale False [-1.1578670740127563, -0.27713173627853394, 1.4349987506866455, -0.03435087203979492, -0.8074522614479065, 0.8418031930923462] grpo gs 3 scale True [-1.0755648612976074, -0.2574329674243927, 1.3329977989196777, -0.05098499357700348, -1.1984542608261108, 1.2494393587112427] rloo [-1.7368006706237793, -0.4156976342201233, 2.152498245239258, -0.05152636766433716, -1.2111784219741821, 1.262704849243164] gmb [-1.1578670740127563, -0.27713173627853394, 1.4349987506866455, -0.03435087203979492, -0.8074522614479065, 0.8418031930923462] [stdout] rewards [-1.8743985891342163, -0.9936632513999939, 0.7184672355651855, -0.6984538435935974, -1.471555233001709, 0.1777002215385437] grpo gs 3 scale False [-1.1578670740127563, -0.27713173627853394, 1.4349987506866455, -0.03435087203979492, -0.8074522614479065, 0.8418031930923462] grpo gs 3 scale True [-1.0755648612976074, -0.2574329674243927, 1.3329977989196777, -0.05098499357700348, -1.1984542608261108, 1.2494393587112427] rloo [-1.7368006706237793, -0.4156976342201233, 2.152498245239258, -0.05152636766433716, -1.2111784219741821, 1.262704849243164] gmb [-1.1578670740127563, -0.27713173627853394, 1.4349987506866455, -0.03435087203979492, -0.8074522614479065, 0.8418031930923462]
Probe gae/lambda_returns
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(7)
rewards=torch.randn(2,4,generator=g)
values=torch.randn(2,4,generator=g)
nextv=torch.randn(2,generator=g)
gamma,lam=0.9,0.95
adv=R.gae(rewards,values,nextv,gamma,lam)
print('gae shape',adv.shape); print(adv.tolist())
print('nextv',nextv.tolist())
# manual gae
T=4
def man_gae(rw,vv,nv):
adv=torch.zeros_like(rw); last=0.0
for t in reversed(range(T)):
nextval = vv[:,t+1] if t+1<T else nv
delta=rw[:,t]+gamma*nextval-vv[:,t]
last=delta+gamma*lam*last
adv[:,t]=last
return adv
print('man',man_gae(rewards,values,nextv).tolist())
print('lambda_returns', R.lambda_returns(rewards,values,nextv,gamma,lam).tolist())
print('man returns', (man_gae(rewards,values,nextv)+values).tolist())
"
Exit code 1 Traceback (most recent call last): File "<string>", line 8, in <module> File "reforacle.py", line 82, in reforacle.gae RuntimeError: The size of tensor a (4) must match the size of tensor b (2) at non-singleton dimension 0 [error] tool reported failure
Probe gae 1D
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(7)
rewards=torch.randn(5,generator=g)
values=torch.randn(5,generator=g)
nextv=torch.randn(1,generator=g).item()*0+0.5
gamma,lam=0.9,0.95
adv=R.gae(rewards,values,torch.tensor(0.5),gamma,lam)
print('gae',adv.shape,adv.tolist())
T=5
def man_gae(rw,vv,nv):
adv=torch.zeros_like(rw); last=0.0
for t in reversed(range(T)):
nextval = vv[t+1] if t+1<T else nv
delta=rw[t]+gamma*nextval-vv[t]
last=delta+gamma*lam*last
adv[t]=last
return adv
print('man',man_gae(rewards,values,0.5).tolist())
print('lr',R.lambda_returns(rewards,values,torch.tensor(0.5),gamma,lam).tolist())
print('man ret',(man_gae(rewards,values,0.5)+values).tolist())
print('disc_ret',R.discounted_returns(rewards,gamma).tolist())
dr=torch.zeros(5); acc=0
for t in reversed(range(5)):
acc=rewards[t]+gamma*acc; dr[t]=acc
print('man dr',dr.tolist())
"
gae torch.Size([5]) [2.5485172271728516, 2.4801719188690186, 0.2681383192539215, 0.5021603107452393, 3.822988748550415] man [2.5485172271728516, 2.4801719188690186, 0.2681383192539215, 0.5021603107452393, 3.822988748550415] lr [1.6536893844604492, 2.124547004699707, 1.5005240440368652, 0.6403329372406006, 2.1407902240753174] man ret [1.6536893844604492, 2.124547004699707, 1.5005240440368652, 0.6403329372406006, 2.1407902240753174] disc_ret [1.6246280670166016, 1.9682480096817017, 1.3134520053863525, 0.40736711025238037, 1.6907901763916016] man dr [1.6246280670166016, 1.9682480096817017, 1.3134520053863525, 0.40736711025238037, 1.6907901763916016] [stdout] gae torch.Size([5]) [2.5485172271728516, 2.4801719188690186, 0.2681383192539215, 0.5021603107452393, 3.822988748550415] man [2.5485172271728516, 2.4801719188690186, 0.2681383192539215, 0.5021603107452393, 3.822988748550415] lr [1.6536893844604492, 2.124547004699707, 1.5005240440368652, 0.6403329372406006, 2.1407902240753174] man ret [1.6536893844604492, 2.124547004699707, 1.5005240440368652, 0.6403329372406006, 2.1407902240753174] disc_ret [1.6246280670166016, 1.9682480096817017, 1.3134520053863525, 0.40736711025238037, 1.6907901763916016] man dr [1.6246280670166016, 1.9682480096817017, 1.3134520053863525, 0.40736711025238037, 1.6907901763916016]
Probe gae time-first 2D
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(7)
rewards=torch.randn(4,2,generator=g)
values=torch.randn(4,2,generator=g)
nextv=torch.randn(2,generator=g)
gamma,lam=0.9,0.95
try:
adv=R.gae(rewards,values,nextv,gamma,lam)
print('gae 4x2',adv.shape,adv.tolist())
except Exception as e: print('err',e)
# time-first manual
T=4
def man(rw,vv,nv):
adv=torch.zeros_like(rw);last=0
for t in reversed(range(T)):
nv_=vv[t+1] if t+1<T else nv
d=rw[t]+gamma*nv_-vv[t]; last=d+gamma*lam*last; adv[t]=last
return adv
print('man',man(rewards,values,nextv).tolist())
"
gae 4x2 torch.Size([4, 2]) [[2.9753334522247314, 1.805307149887085], [3.4788193702697754, -0.9152735471725464], [3.1884028911590576, 0.1349542737007141], [0.44331133365631104, 0.06784558296203613]] man [[2.9753334522247314, 1.805307149887085], [3.4788193702697754, -0.9152735471725464], [3.1884028911590576, 0.1349542737007141], [0.44331133365631104, 0.06784558296203613]] [stdout] gae 4x2 torch.Size([4, 2]) [[2.9753334522247314, 1.805307149887085], [3.4788193702697754, -0.9152735471725464], [3.1884028911590576, 0.1349542737007141], [0.44331133365631104, 0.06784558296203613]] man [[2.9753334522247314, 1.805307149887085], [3.4788193702697754, -0.9152735471725464], [3.1884028911590576, 0.1349542737007141], [0.44331133365631104, 0.06784558296203613]]
Probe kl_penalty estimators
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(8)
logp=torch.randn(5,generator=g); ref=torch.randn(5,generator=g)
for est in ['k1','k2','k3','kl','mse','abs']:
try: print('kl',est, R.kl_penalty(logp,ref,est).tolist())
except Exception as e: print('kl',est,'ERR',str(e)[:60])
d=logp-ref
print('k1=logp-ref', d.tolist())
print('k2=.5 d^2',(0.5*d*d).tolist())
print('k3=exp(-d)-1+d? ', (torch.exp(ref-logp)-1-(ref-logp)).tolist())
print('rev', R.reverse_kl(logp,ref))
print('sym', R.symmetric_kl(logp,ref))
"
kl k1 [0.9320472478866577, -1.1296039819717407, 0.6070117950439453, 1.9493353366851807, -1.27122163772583] kl k2 [0.43435603380203247, 0.6380025744438171, 0.18423165380954742, 1.899954080581665, 0.808002233505249] kl k3 [0.325793981552124, 0.9648269414901733, 0.15198874473571777, 1.0917038917541504, 1.2939836978912354] kl kl ERR kl kl mse ERR mse kl abs ERR abs k1=logp-ref [0.9320472478866577, -1.1296039819717407, 0.6070117950439453, 1.9493353366851807, -1.27122163772583] k2=.5 d^2 [0.43435603380203247, 0.6380025744438171, 0.18423165380954742, 1.899954080581665, 0.808002233505249] k3=exp(-d)-1+d? [0.3257940411567688, 0.9648269414901733, 0.15198874473571777, 1.09170401096344, 1.2939836978912354] rev tensor([0.3258, 0.9648, 0.1520, 1.0917, 1.2940]) sym tensor([0.4667, 0.7088, 0.1900, 2.5832, 0.9228]) [stdout] kl k1 [0.9320472478866577, -1.1296039819717407, 0.6070117950439453, 1.9493353366851807, -1.27122163772583] kl k2 [0.43435603380203247, 0.6380025744438171, 0.18423165380954742, 1.899954080581665, 0.808002233505249] kl k3 [0.325793981552124, 0.9648269414901733, 0.15198874473571777, 1.0917038917541504, 1.2939836978912354] kl kl ERR kl kl mse ERR mse kl abs ERR abs k1=logp-ref [0.9320472478866577, -1.1296039819717407, 0.6070117950439453, 1.9493353366851807, -1.27122163772583] k2=.5 d^2 [0.43435603380203247, 0.6380025744438171, 0.18423165380954742, 1.899954080581665, 0.808002233505249] k3=exp(-d)-1+d? [0.3257940411567688, 0.9648269414901733, 0.15198874473571777, 1.09170401096344, 1.2939836978912354] rev tensor([0.3258, 0.9648, 0.1520, 1.0917, 1.2940]) sym tensor([0.4667, 0.7088, 0.1900, 2.5832, 0.9228])
Probe symmetric_kl
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(8)
logp=torch.randn(5,generator=g); ref=torch.randn(5,generator=g)
d=logp-ref
sym=R.symmetric_kl(logp,ref)
rev=R.reverse_kl(logp,ref)
fwd=torch.exp(d)-1-d # k3 with roles swapped
print('sym',sym.tolist())
print('fwd(exp(d)-1-d)',fwd.tolist())
print('rev',rev.tolist())
print('fwd+rev',(fwd+rev).tolist())
print('.5(fwd+rev)',(0.5*(fwd+rev)).tolist())
"
sym [0.46672505140304565, 0.7087960839271545, 0.18995845317840576, 2.583193063735962, 0.9228470325469971] fwd(exp(d)-1-d) [0.6076561212539673, 0.45276522636413574, 0.22792816162109375, 4.074682235717773, 0.5517103672027588] rev [0.325793981552124, 0.9648269414901733, 0.15198874473571777, 1.0917038917541504, 1.2939836978912354] fwd+rev [0.9334501028060913, 1.417592167854309, 0.3799169063568115, 5.166386127471924, 1.8456940650939941] .5(fwd+rev) [0.46672505140304565, 0.7087960839271545, 0.18995845317840576, 2.583193063735962, 0.9228470325469971] [stdout] sym [0.46672505140304565, 0.7087960839271545, 0.18995845317840576, 2.583193063735962, 0.9228470325469971] fwd(exp(d)-1-d) [0.6076561212539673, 0.45276522636413574, 0.22792816162109375, 4.074682235717773, 0.5517103672027588] rev [0.325793981552124, 0.9648269414901733, 0.15198874473571777, 1.0917038917541504, 1.2939836978912354] fwd+rev [0.9334501028060913, 1.417592167854309, 0.3799169063568115, 5.166386127471924, 1.8456940650939941] .5(fwd+rev) [0.46672505140304565, 0.7087960839271545, 0.18995845317840576, 2.583193063735962, 0.9228470325469971]
Probe importance_ratio/clip_fraction
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(9)
logp=torch.randn(5,generator=g); old=torch.randn(5,generator=g)
print('imp clip None',R.importance_ratio(logp,old,None).tolist())
print('exp',torch.exp(logp-old).tolist())
print('imp clip 0.2',R.importance_ratio(logp,old,0.2).tolist())
r=torch.exp(logp-old)
print('clamp',torch.clamp(r,1-0.2,1+0.2).tolist())
# clip_fraction
print('clip_fraction',R.clip_fraction(logp,old,0.2))
print('frac calc',((r-1).abs()>0.2).float().mean())
"
imp clip None [0.7225927710533142, 4.6049299240112305, 0.5897426009178162, 0.6694101095199585, 1.259422779083252] exp [0.7225927710533142, 4.6049299240112305, 0.5897426009178162, 0.6694101095199585, 1.259422779083252] imp clip 0.2 [0.800000011920929, 1.2000000476837158, 0.800000011920929, 0.800000011920929, 1.2000000476837158] clamp [0.800000011920929, 1.2000000476837158, 0.800000011920929, 0.800000011920929, 1.2000000476837158] clip_fraction tensor(1.) frac calc tensor(1.) [stdout] imp clip None [0.7225927710533142, 4.6049299240112305, 0.5897426009178162, 0.6694101095199585, 1.259422779083252] exp [0.7225927710533142, 4.6049299240112305, 0.5897426009178162, 0.6694101095199585, 1.259422779083252] imp clip 0.2 [0.800000011920929, 1.2000000476837158, 0.800000011920929, 0.800000011920929, 1.2000000476837158] clamp [0.800000011920929, 1.2000000476837158, 0.800000011920929, 0.800000011920929, 1.2000000476837158] clip_fraction tensor(1.) frac calc tensor(1.)
Probe clip_fraction mixed
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(11)
logp=torch.randn(20,generator=g)*0.2; old=torch.randn(20,generator=g)*0.2
r=torch.exp(logp-old)
print('cf 0.2',R.clip_fraction(logp,old,0.2).item())
print('abs>c',((r-1).abs()>0.2).float().mean().item())
print('gt',(r-1).abs().gt(0.2).float().mean().item())
"
cf 0.2 0.3499999940395355 abs>c 0.3499999940395355 gt 0.3499999940395355 [stdout] cf 0.2 0.3499999940395355 abs>c 0.3499999940395355 gt 0.3499999940395355
Probe clipped_pg_loss
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(12)
logp=torch.randn(2,4,generator=g); old=torch.randn(2,4,generator=g)
adv=torch.randn(2,4,generator=g); mask=(torch.rand(2,4,generator=g)>0.3).float()
cl=0.2;ch=0.2
loss=R.clipped_pg_loss(logp,old,adv,mask,cl,ch)
print('cpg',loss.item())
r=torch.exp(logp-old)
unc=r*adv
cl_=torch.clamp(r,1-cl,1+ch)*adv
pg=-torch.min(unc,cl_)
print('masked mean',((pg*mask).sum()/mask.sum()).item())
print('mean all',pg.mean().item())
"
cpg -0.06370402872562408 masked mean -0.06370402872562408 mean all -0.02508426457643509 [stdout] cpg -0.06370402872562408 masked mean -0.06370402872562408 mean all -0.02508426457643509
Probe value_loss/huber
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(13)
v=torch.randn(2,4,generator=g); ov=torch.randn(2,4,generator=g); ret=torch.randn(2,4,generator=g)
print('value_loss clip0.2',R.value_loss(v,ov,ret,0.2).item())
vc=ov+torch.clamp(v-ov,-0.2,0.2)
l1=(v-ret)**2; l2=(vc-ret)**2
print('.5 mean max',(0.5*torch.max(l1,l2)).mean().item())
print('mean max nofactor',(torch.max(l1,l2)).mean().item())
print('value_loss clipNone',R.value_loss(v,ov,ret,None).item())
print('.5 mean (v-ret)^2',(0.5*l1).mean().item())
print('huber',R.huber_value_loss(v,ret,1.0).item())
import torch.nn.functional as F
print('F huber',F.huber_loss(v,ret,delta=1.0).item())
print('F smoothl1',F.smooth_l1_loss(v,ret).item())
"
Exit code 1 value_loss clip0.2 1.1638469696044922 .5 mean max 1.1638469696044922 mean max nofactor 2.3276939392089844 Traceback (most recent call last): File "<string>", line 10, in <module> File "reforacle.py", line 109, in reforacle.value_loss TypeError: bad operand type for unary -: 'NoneType' [error] tool reported failure
Probe huber
python3 -c "
import torch, reforacle as R
import torch.nn.functional as F
g=torch.Generator().manual_seed(13)
v=torch.randn(2,4,generator=g); ret=torch.randn(2,4,generator=g)
print('huber d1',R.huber_value_loss(v,ret,1.0).item())
print('F huber d1',F.huber_loss(v,ret,delta=1.0).item())
print('huber d0.5',R.huber_value_loss(v,ret,0.5).item())
print('F huber d0.5',F.huber_loss(v,ret,delta=0.5).item())
"
huber d1 0.7040889859199524 F huber d1 0.7040889859199524 huber d0.5 0.450337678194046 F huber d0.5 0.450337678194046 [stdout] huber d1 0.7040889859199524 F huber d1 0.7040889859199524 huber d0.5 0.450337678194046 F huber d0.5 0.450337678194046
Probe whiten
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(14)
v=torch.randn(8,generator=g)
mask=(torch.rand(8,generator=g)>0.3).float()
for sm in [True,False]:
print('whiten sm',sm,R.whiten(v,mask,sm).tolist())
# whiten ignoring mask? or using it. Let me compute unmasked whiten
mean=v.mean(); std=v.std() # unbiased
print('unbiased',( (v-mean)/(std+1e-8) ).tolist())
mean=v.mean(); var=v.var(unbiased=False)
print('biased',((v-mean)/(torch.sqrt(var)+1e-8)).tolist())
print('masked_whiten sm True', R.masked_whiten(v,mask,True).tolist())
"
whiten sm True [-0.8619763255119324, -0.1025102362036705, -0.5514551997184753, -0.682371973991394, -0.049288056790828705, 0.9539925456047058, 1.8304486274719238, -1.3988155126571655] whiten sm False [-1.1473407745361328, -0.38787466287612915, -0.8368196487426758, -0.9677363634109497, -0.33465248346328735, 0.6686280965805054, 1.5450842380523682, -1.6841799020767212] unbiased [-0.7214695811271667, 0.005009348504245281, -0.4244358241558075, -0.5496663451194763, 0.0559198372066021, 1.0156233310699463, 1.854010820388794, -1.2349913120269775] biased [-0.7712835073471069, 0.005355218891054392, -0.45374101400375366, -0.5876180529594421, 0.059780824929475784, 1.0857470035552979, 1.9820209741592407, -1.3202613592147827] masked_whiten sm True [-0.0, -0.1025102362036705, -0.5514551997184753, -0.682371973991394, -0.049288056790828705, 0.9539925456047058, 1.8304486274719238, -1.3988155126571655] [stdout] whiten sm True [-0.8619763255119324, -0.1025102362036705, -0.5514551997184753, -0.682371973991394, -0.049288056790828705, 0.9539925456047058, 1.8304486274719238, -1.3988155126571655] whiten sm False [-1.1473407745361328, -0.38787466287612915, -0.8368196487426758, -0.9677363634109497, -0.33465248346328735, 0.6686280965805054, 1.5450842380523682, -1.6841799020767212] unbiased [-0.7214695811271667, 0.005009348504245281, -0.4244358241558075, -0.5496663451194763, 0.0559198372066021, 1.0156233310699463, 1.854010820388794, -1.2349913120269775] biased [-0.7712835073471069, 0.005355218891054392, -0.45374101400375366, -0.5876180529594421, 0.059780824929475784, 1.0857470035552979, 1.9820209741592407, -1.3202613592147827] masked_whiten sm True [-0.0, -0.1025102362036705, -0.5514551997184753, -0.682371973991394, -0.049288056790828705, 0.9539925456047058, 1.8304486274719238, -1.3988155126571655]
Probe whiten masked stats
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(14)
v=torch.randn(8,generator=g)
mask=(torch.rand(8,generator=g)>0.3).float()
print('mask',mask.tolist())
n=mask.sum()
mean=(v*mask).sum()/n
var=((v-mean)**2*mask).sum()/n # biased masked
varub=((v-mean)**2*mask).sum()/(n-1)
w=R.whiten(v,mask,True)
print('ref',w.tolist())
print('biased',((v-mean)/torch.sqrt(var+1e-8)).tolist())
print('ubias',((v-mean)/torch.sqrt(varub+1e-8)).tolist())
print('ubias noeps',((v-mean)/torch.sqrt(varub)).tolist())
"
mask [0.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0] ref [-0.8619763255119324, -0.1025102362036705, -0.5514551997184753, -0.682371973991394, -0.049288056790828705, 0.9539925456047058, 1.8304486274719238, -1.3988155126571655] biased [-0.8619763255119324, -0.1025102362036705, -0.5514551997184753, -0.682371973991394, -0.049288056790828705, 0.9539925456047058, 1.8304486274719238, -1.3988155126571655] ubias [-0.7980349659919739, -0.0949060320854187, -0.5105482339859009, -0.63175368309021, -0.04563187062740326, 0.8832253813743591, 1.6946660280227661, -1.2950514554977417] ubias noeps [-0.7980349659919739, -0.0949060320854187, -0.5105482339859009, -0.63175368309021, -0.04563187062740326, 0.8832253813743591, 1.6946660280227661, -1.2950514554977417] [stdout] mask [0.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0] ref [-0.8619763255119324, -0.1025102362036705, -0.5514551997184753, -0.682371973991394, -0.049288056790828705, 0.9539925456047058, 1.8304486274719238, -1.3988155126571655] biased [-0.8619763255119324, -0.1025102362036705, -0.5514551997184753, -0.682371973991394, -0.049288056790828705, 0.9539925456047058, 1.8304486274719238, -1.3988155126571655] ubias [-0.7980349659919739, -0.0949060320854187, -0.5105482339859009, -0.63175368309021, -0.04563187062740326, 0.8832253813743591, 1.6946660280227661, -1.2950514554977417] ubias noeps [-0.7980349659919739, -0.0949060320854187, -0.5105482339859009, -0.63175368309021, -0.04563187062740326, 0.8832253813743591, 1.6946660280227661, -1.2950514554977417]
Confirm whiten shift_mean False
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(14)
v=torch.randn(8,generator=g)
mask=(torch.rand(8,generator=g)>0.3).float()
n=mask.sum(); mean=(v*mask).sum()/n
var=((v-mean)**2*mask).sum()/n
wh=(v-mean)*torch.rsqrt(var+1e-8)
print('smFalse ref',R.whiten(v,mask,False).tolist())
print('wh+mean',(wh+mean).tolist())
# masked_whiten sm False
print('mwh smFalse ref',R.masked_whiten(v,mask,False).tolist())
print('(wh+mean)*mask',((wh+mean)*mask).tolist())
"
smFalse ref [-1.1473407745361328, -0.38787466287612915, -0.8368196487426758, -0.9677363634109497, -0.33465248346328735, 0.6686280965805054, 1.5450842380523682, -1.6841799020767212] wh+mean [-1.1473407745361328, -0.38787466287612915, -0.8368195295333862, -0.9677363634109497, -0.33465248346328735, 0.6686280965805054, 1.5450841188430786, -1.6841799020767212] mwh smFalse ref [-0.28536441922187805, -0.38787466287612915, -0.8368196487426758, -0.9677363634109497, -0.33465248346328735, 0.6686280965805054, 1.5450842380523682, -1.6841799020767212] (wh+mean)*mask [-0.0, -0.38787466287612915, -0.8368195295333862, -0.9677363634109497, -0.33465248346328735, 0.6686280965805054, 1.5450841188430786, -1.6841799020767212] [stdout] smFalse ref [-1.1473407745361328, -0.38787466287612915, -0.8368196487426758, -0.9677363634109497, -0.33465248346328735, 0.6686280965805054, 1.5450842380523682, -1.6841799020767212] wh+mean [-1.1473407745361328, -0.38787466287612915, -0.8368195295333862, -0.9677363634109497, -0.33465248346328735, 0.6686280965805054, 1.5450841188430786, -1.6841799020767212] mwh smFalse ref [-0.28536441922187805, -0.38787466287612915, -0.8368196487426758, -0.9677363634109497, -0.33465248346328735, 0.6686280965805054, 1.5450842380523682, -1.6841799020767212] (wh+mean)*mask [-0.0, -0.38787466287612915, -0.8368195295333862, -0.9677363634109497, -0.33465248346328735, 0.6686280965805054, 1.5450841188430786, -1.6841799020767212]
Probe normalize
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(15)
x=torch.randn(8,generator=g)
print('normalize eps1e-8',R.normalize(x,1e-8).tolist())
m=x.mean();s=x.std()
print('unbiased',((x-m)/(s+1e-8)).tolist())
sb=x.var(unbiased=False).sqrt()
print('biased',((x-m)/(sb+1e-8)).tolist())
print('rsqrt biased',((x-m)*torch.rsqrt(x.var(unbiased=False)+1e-8)).tolist())
"
normalize eps1e-8 [-0.8574098944664001, 0.7921727895736694, -0.6657798886299133, 1.0750460624694824, 1.2633132934570312, 0.156465545296669, -1.5385518074035645, -0.225256085395813] unbiased [-0.8574098944664001, 0.7921727895736694, -0.6657798886299133, 1.0750460624694824, 1.2633132934570312, 0.156465545296669, -1.5385518074035645, -0.225256085395813] biased [-0.9166097640991211, 0.8468683362007141, -0.7117486596107483, 1.1492726802825928, 1.350538730621338, 0.16726870834827423, -1.6447809934616089, -0.2408088892698288] rsqrt biased [-0.9166097640991211, 0.8468683362007141, -0.7117486596107483, 1.1492726802825928, 1.350538730621338, 0.16726870834827423, -1.6447809934616089, -0.2408088892698288] [stdout] normalize eps1e-8 [-0.8574098944664001, 0.7921727895736694, -0.6657798886299133, 1.0750460624694824, 1.2633132934570312, 0.156465545296669, -1.5385518074035645, -0.225256085395813] unbiased [-0.8574098944664001, 0.7921727895736694, -0.6657798886299133, 1.0750460624694824, 1.2633132934570312, 0.156465545296669, -1.5385518074035645, -0.225256085395813] biased [-0.9166097640991211, 0.8468683362007141, -0.7117486596107483, 1.1492726802825928, 1.350538730621338, 0.16726870834827423, -1.6447809934616089, -0.2408088892698288] rsqrt biased [-0.9166097640991211, 0.8468683362007141, -0.7117486596107483, 1.1492726802825928, 1.350538730621338, 0.16726870834827423, -1.6447809934616089, -0.2408088892698288]
Probe smoothed_nll/cross_entropy
python3 -c "
import torch, reforacle as R
import torch.nn.functional as F
g=torch.Generator().manual_seed(16)
logits=torch.randn(2,4,6,generator=g)
labels=torch.randint(0,6,(2,4),generator=g)
print('smoothed_nll s0.1',R.smoothed_nll(logits,labels,0.1).item())
lsm=torch.log_softmax(logits,-1)
nll=-lsm.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
smooth=-lsm.mean(-1)
sn=(1-0.1)*nll+0.1*smooth
print('formula mean',sn.mean().item())
print('cross_entropy ignore -100',R.cross_entropy(logits,labels,-100).item())
print('F ce',F.cross_entropy(logits.reshape(-1,6),labels.reshape(-1),ignore_index=-100).item())
print('smoothed_nll s0',R.smoothed_nll(logits,labels,0.0).item())
print('nll mean',nll.mean().item())
"
smoothed_nll s0.1 2.2135331630706787 formula mean 2.2135331630706787 cross_entropy ignore -100 2.212557792663574 F ce 2.2125580310821533 smoothed_nll s0 2.212557792663574 nll mean 2.212557792663574 [stdout] smoothed_nll s0.1 2.2135331630706787 formula mean 2.2135331630706787 cross_entropy ignore -100 2.212557792663574 F ce 2.2125580310821533 smoothed_nll s0 2.212557792663574 nll mean 2.212557792663574
Probe cross_entropy ignore/bradley_terry
python3 -c "
import torch, reforacle as R
import torch.nn.functional as F
g=torch.Generator().manual_seed(17)
logits=torch.randn(3,5,7,generator=g)
labels=torch.randint(0,7,(3,5),generator=g)
labels[0,0]=-100; labels[1,2]=-100
print('ce',R.cross_entropy(logits,labels,-100).item())
print('F',F.cross_entropy(logits.reshape(-1,7),labels.reshape(-1),ignore_index=-100).item())
# smoothed_nll with ignore? does it accept? probably no ignore
# bradley_terry_logit
print('bt',R.bradley_terry_logit(torch.tensor(1.0),torch.tensor(0.3),0.5))
print('sigmoid(beta*(c-r))',torch.sigmoid(torch.tensor(0.5*(1.0-0.3))))
print('beta*(c-r)',0.5*(1.0-0.3))
"
ce 2.299717903137207 F 2.299718141555786 bt tensor(0.3500) sigmoid(beta*(c-r)) tensor(0.5866) beta*(c-r) 0.35 [stdout] ce 2.299717903137207 F 2.299718141555786 bt tensor(0.3500) sigmoid(beta*(c-r)) tensor(0.5866) beta*(c-r) 0.35
Probe argmax/mode/median
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(18)
logits=torch.randn(2,4,6,generator=g)
at=R.argmax_tokens(logits)
print('argmax',at.shape,at.dtype,at.tolist())
print('torch',logits.argmax(-1).tolist())
labels=torch.tensor([[1,2,2,3],[0,0,1,1]])
print('mode_label',R.mode_label(labels), R.mode_label(labels).shape if hasattr(R.mode_label(labels),'shape') else '')
rewards=torch.tensor([3.0,1.0,2.0,5.0])
print('median',R.median_reward(rewards).item(), rewards.median().item())
rw5=torch.tensor([3.0,1.0,2.0,5.0,4.0])
print('median5',R.median_reward(rw5).item())
"
argmax torch.Size([2, 4]) torch.int64 [[3, 3, 1, 1], [1, 0, 4, 2]] torch [[3, 3, 1, 1], [1, 0, 4, 2]] mode_label tensor([2, 1]) torch.Size([2]) median 2.5 2.0 median5 3.0 [stdout] argmax torch.Size([2, 4]) torch.int64 [[3, 3, 1, 1], [1, 0, 4, 2]] torch [[3, 3, 1, 1], [1, 0, 4, 2]] mode_label tensor([2, 1]) torch.Size([2]) median 2.5 2.0 median5 3.0
Probe mode/median details
python3 -c "
import torch, reforacle as R
print('torch.mode [0,0,1,1]',torch.mode(torch.tensor([0,0,1,1])).values.item())
rewards=torch.tensor([3.0,1.0,2.0,5.0])
print('quantile0.5',torch.quantile(rewards,0.5).item())
labels=torch.tensor([[1,2,2,3],[0,0,1,1]])
print('torch.mode dim-1',torch.mode(labels,dim=-1).values.tolist())
# median_reward on 2d?
r2=torch.tensor([[3.0,1.0,2.0,5.0],[1.,1.,1.,1.]])
try: print('median 2d',R.median_reward(r2))
except Exception as e: print('err',e)
"
torch.mode [0,0,1,1] 0 quantile0.5 2.5 torch.mode dim-1 [2, 0] median 2d tensor([1.0000, 1.5000, 2.0000, 3.0000]) [stdout] torch.mode [0,0,1,1] 0 quantile0.5 2.5 torch.mode dim-1 [2, 0] median 2d tensor([1.0000, 1.5000, 2.0000, 3.0000])
Probe median_reward 2d
python3 -c "
import torch, reforacle as R
r2=torch.tensor([[3.0,1.0,2.0,5.0],[1.,1.,1.,1.]])
print('q dim0',torch.quantile(r2,0.5,dim=0).tolist())
print('q dim1',torch.quantile(r2,0.5,dim=1).tolist())
print('median dim0',torch.median(r2,dim=0).values.tolist())
print('ref',R.median_reward(r2).tolist())
print('ref 1d row0',R.median_reward(r2[0]).item())
# maybe quantile flatten sorted?
print('sorted quantile of all?', torch.quantile(r2.flatten(),0.5).item())
"
q dim0 [2.0, 1.0, 1.5, 3.0] q dim1 [2.5, 1.0] median dim0 [1.0, 1.0, 1.0, 1.0] ref [1.0, 1.5, 2.0, 3.0] ref 1d row0 2.5 sorted quantile of all? 1.0 [stdout] q dim0 [2.0, 1.0, 1.5, 3.0] q dim1 [2.5, 1.0] median dim0 [1.0, 1.0, 1.0, 1.0] ref [1.0, 1.5, 2.0, 3.0] ref 1d row0 2.5 sorted quantile of all? 1.0
Probe median_reward pattern
python3 -c "
import torch, reforacle as R
a=torch.tensor([[5.0,4.0,3.0,2.0,1.0],[10.,20.,30.,40.,50.]])
print('ref',R.median_reward(a).tolist())
print('median all',a.median().item())
print('q0.5 dim0',torch.quantile(a,0.5,dim=0).tolist())
print('q0.5 dim1',torch.quantile(a,0.5,dim=1).tolist())
b=torch.arange(6.).reshape(2,3)
print('b ref',R.median_reward(b).tolist(), 'shape')
"
ref [5.5, 11.0, 16.5, 22.0, 27.5] median all 5.0 q0.5 dim0 [7.5, 12.0, 16.5, 21.0, 25.5] q0.5 dim1 [3.0, 30.0] b ref [1.5, 2.5, 3.5] shape [stdout] ref [5.5, 11.0, 16.5, 22.0, 27.5] median all 5.0 q0.5 dim0 [7.5, 12.0, 16.5, 21.0, 25.5] q0.5 dim1 [3.0, 30.0] b ref [1.5, 2.5, 3.5] shape
Verify median_reward hypothesis
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(99)
a=torch.randn(3,4,generator=g)
ref=R.median_reward(a)
s=a.flatten().sort().values.reshape(a.shape)
mine=torch.quantile(s,0.5,dim=0)
print('ref',ref.tolist())
print('mine',mine.tolist())
print('match',torch.allclose(ref,mine))
# also 1d
b=torch.randn(7,generator=g)
print('1d',R.median_reward(b).item(), torch.quantile(b.flatten().sort().values,0.5,dim=0).item(), b.median().item())
"
ref [-1.3890278339385986, -0.6453173756599426, -0.2729676067829132, 0.7443659901618958] mine [-0.7646492719650269, -0.6665656566619873, -0.6453173756599426, -0.2729676067829132] match False 1d -0.027061855420470238 -0.027061855420470238 -0.027061855420470238 [stdout] ref [-1.3890278339385986, -0.6453173756599426, -0.2729676067829132, 0.7443659901618958] mine [-0.7646492719650269, -0.6665656566619873, -0.6453173756599426, -0.2729676067829132] match False 1d -0.027061855420470238 -0.027061855420470238 -0.027061855420470238
Brute median_reward 3x4
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(99)
a=torch.randn(3,4,generator=g)
ref=R.median_reward(a)
print('ref',ref.tolist())
print('q dim0',torch.quantile(a,0.5,dim=0).tolist())
print('q dim1',torch.quantile(a,0.5,dim=1).tolist())
# sort each column (dim0) then quantile? same as quantile dim0
# sort each row then quantile dim0
sr=a.sort(dim=1).values
print('rowsort q dim0',torch.quantile(sr,0.5,dim=0).tolist())
sc=a.sort(dim=0).values
print('colsort q dim1',torch.quantile(sc,0.5,dim=1).tolist())
print('colsort q dim0',torch.quantile(sc,0.5,dim=0).tolist())
print('a',a.tolist())
"
ref [-1.3890278339385986, -0.6453173756599426, -0.2729676067829132, 0.7443659901618958] q dim0 [0.7443659901618958, -1.1753536462783813, -0.9511449337005615, -0.2729676067829132] q dim1 [-0.7156074643135071, -0.4591425061225891, -0.5853967070579529] rowsort q dim0 [-1.3890278339385986, -0.7646492719650269, -0.2729676067829132, 0.7443659901618958] colsort q dim1 [-1.027796745300293, -0.6120562553405762, -0.4324829578399658] colsort q dim0 [0.7443659901618958, -1.1753536462783813, -0.9511449337005615, -0.2729676067829132] a [[0.6126858592033386, -1.1753536462783813, -0.7646492719650269, -0.6665656566619873], [0.7443659901618958, -0.6453173756599426, -1.3890278339385986, -0.2729676067829132], [0.940598726272583, -2.6176693439483643, -0.9511449337005615, -0.21964849531650543]] [stdout] ref [-1.3890278339385986, -0.6453173756599426, -0.2729676067829132, 0.7443659901618958] q dim0 [0.7443659901618958, -1.1753536462783813, -0.9511449337005615, -0.2729676067829132] q dim1 [-0.7156074643135071, -0.4591425061225891, -0.5853967070579529] rowsort q dim0 [-1.3890278339385986, -0.7646492719650269, -0.2729676067829132, 0.7443659901618958] colsort q dim1 [-1.027796745300293, -0.6120562553405762, -0.4324829578399658] colsort q dim0 [0.7443659901618958, -1.1753536462783813, -0.9511449337005615, -0.2729676067829132] a [[0.6126858592033386, -1.1753536462783813, -0.7646492719650269, -0.6665656566619873], [0.7443659901618958, -0.6453173756599426, -1.3890278339385986, -0.2729676067829132], [0.940598726272583, -2.6176693439483643, -0.9511449337005615, -0.21964849531650543]]
flat sort reshape 3x4
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(99)
a=torch.randn(3,4,generator=g)
ref=R.median_reward(a)
fs=a.flatten().sort().values.reshape(3,4)
print('flat-sort-reshape q dim0',torch.quantile(fs,0.5,dim=0).tolist())
print('ref',ref.tolist())
print('fs',fs.tolist())
# row1 sorted
print('row1 sorted',a[1].sort().values.tolist())
"
flat-sort-reshape q dim0 [-0.7646492719650269, -0.6665656566619873, -0.6453173756599426, -0.2729676067829132] ref [-1.3890278339385986, -0.6453173756599426, -0.2729676067829132, 0.7443659901618958] fs [[-2.6176693439483643, -1.3890278339385986, -1.1753536462783813, -0.9511449337005615], [-0.7646492719650269, -0.6665656566619873, -0.6453173756599426, -0.2729676067829132], [-0.21964849531650543, 0.6126858592033386, 0.7443659901618958, 0.940598726272583]] row1 sorted [-1.3890278339385986, -0.6453173756599426, -0.2729676067829132, 0.7443659901618958] [stdout] flat-sort-reshape q dim0 [-0.7646492719650269, -0.6665656566619873, -0.6453173756599426, -0.2729676067829132] ref [-1.3890278339385986, -0.6453173756599426, -0.2729676067829132, 0.7443659901618958] fs [[-2.6176693439483643, -1.3890278339385986, -1.1753536462783813, -0.9511449337005615], [-0.7646492719650269, -0.6665656566619873, -0.6453173756599426, -0.2729676067829132], [-0.21964849531650543, 0.6126858592033386, 0.7443659901618958, 0.940598726272583]] row1 sorted [-1.3890278339385986, -0.6453173756599426, -0.2729676067829132, 0.7443659901618958]
Decode median_reward mapping
python3 -c "
import torch, reforacle as R
a=torch.tensor([[10.,11,12,13],[20,21,22,23],[30,31,32,33]])
print('ref',R.median_reward(a).tolist())
a2=torch.tensor([[30.,10,20,40],[3,1,2,4],[300,100,200,400]])
print('ref2',R.median_reward(a2).tolist())
# guess: sort flatten, reshape, take middle row
fs=a2.flatten().sort().values.reshape(3,4)
print('fs2',fs.tolist())
print('fs2 mid row',fs[1].tolist())
"
ref [20.0, 21.0, 22.0, 23.0] ref2 [1.0, 2.0, 3.0, 4.0] fs2 [[1.0, 2.0, 3.0, 4.0], [10.0, 20.0, 30.0, 40.0], [100.0, 200.0, 300.0, 400.0]] fs2 mid row [10.0, 20.0, 30.0, 40.0] [stdout] ref [20.0, 21.0, 22.0, 23.0] ref2 [1.0, 2.0, 3.0, 4.0] fs2 [[1.0, 2.0, 3.0, 4.0], [10.0, 20.0, 30.0, 40.0], [100.0, 200.0, 300.0, 400.0]] fs2 mid row [10.0, 20.0, 30.0, 40.0]
Test median_reward atomic-row hypothesis
python3 -c "
import torch, reforacle as R, math
def mine(rewards):
sr=rewards.sort(dim=-1).values if rewards.dim()>0 else rewards
if rewards.dim()==0: return rewards
n=sr.shape[0]; pos=0.5*(n-1)
lo=math.floor(pos); hi=math.ceil(pos); frac=pos-lo
return sr[lo]*(1-frac)+sr[hi]*frac
for shape in [(5,),(3,4),(4,5),(6,3),(2,7)]:
g=torch.Generator().manual_seed(sum(shape)+7)
a=torch.randn(*shape,generator=g)
ref=R.median_reward(a)
try:
m=mine(a)
print(shape,'match',torch.allclose(ref,m,atol=1e-5), ref.shape, m.shape)
except Exception as e:
print(shape,'ERR',e, 'ref',ref.tolist() if hasattr(ref,'tolist') else ref)
"
(5,) match True torch.Size([]) torch.Size([]) (3, 4) match True torch.Size([4]) torch.Size([4]) (4, 5) match True torch.Size([5]) torch.Size([5]) (6, 3) match True torch.Size([3]) torch.Size([3]) (2, 7) match True torch.Size([7]) torch.Size([7]) [stdout] (5,) match True torch.Size([]) torch.Size([]) (3, 4) match True torch.Size([4]) torch.Size([4]) (4, 5) match True torch.Size([5]) torch.Size([5]) (6, 3) match True torch.Size([3]) torch.Size([3]) (2, 7) match True torch.Size([7]) torch.Size([7])
Probe masked_max/quantile_lower/advantage_mean_std
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(20)
# masked_max
v=torch.randn(2,4,generator=g); m=(torch.rand(2,4,generator=g)>0.4).float()
print('masked_max',R.masked_max(v,m))
print('v',v.tolist());print('m',m.tolist())
# quantile_lower
x=torch.randn(10,generator=g)
print('ql 0.25',R.quantile_lower(x,0.25).item(), torch.quantile(x,0.25).item())
print('ql lower interp', torch.quantile(x,0.25,interpolation='lower').item())
# advantage_mean_std
adv=torch.randn(2,4,generator=g); am=(torch.rand(2,4,generator=g)>0.3).float()
print('ams',R.advantage_mean_std(adv,am))
"
masked_max tensor(1.1632) v [[-1.206070899963379, 0.06167072802782059, 1.1631553173065186, -1.5008302927017212], [-1.5944066047668457, -0.01870560459792614, -2.1324880123138428, -0.5270004868507385]] m [[1.0, 1.0, 1.0, 0.0], [1.0, 0.0, 1.0, 0.0]] ql 0.25 -0.613038182258606 -0.5822517275810242 ql lower interp -0.613038182258606 ams tensor([0.0783, 0.8284]) [stdout] masked_max tensor(1.1632) v [[-1.206070899963379, 0.06167072802782059, 1.1631553173065186, -1.5008302927017212], [-1.5944066047668457, -0.01870560459792614, -2.1324880123138428, -0.5270004868507385]] m [[1.0, 1.0, 1.0, 0.0], [1.0, 0.0, 1.0, 0.0]] ql 0.25 -0.613038182258606 -0.5822517275810242 ql lower interp -0.613038182258606 ams tensor([0.0783, 0.8284])
Probe advantage_mean_std std type
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(20)
v=torch.randn(2,4,generator=g); m=(torch.rand(2,4,generator=g)>0.4).float()
x=torch.randn(10,generator=g)
adv=torch.randn(2,4,generator=g); am=(torch.rand(2,4,generator=g)>0.3).float()
n=am.sum(); mean=(adv*am).sum()/n
varb=((adv-mean)**2*am).sum()/n
varu=((adv-mean)**2*am).sum()/(n-1)
print('ref',R.advantage_mean_std(adv,am).tolist())
print('mean,biased std',[mean.item(),varb.sqrt().item()])
print('mean,unbiased std',[mean.item(),varu.sqrt().item()])
# masked_max along axis? test 2d with axis
print('masked_max',R.masked_max(v,m).shape)
"
ref [0.07827985286712646, 0.8283936977386475] mean,biased std [0.07827985286712646, 0.8283936977386475] mean,unbiased std [0.07827985286712646, 0.8947674632072449] masked_max torch.Size([]) [stdout] ref [0.07827985286712646, 0.8283936977386475] mean,biased std [0.07827985286712646, 0.8283936977386475] mean,unbiased std [0.07827985286712646, 0.8947674632072449] masked_max torch.Size([])
Probe top_k/top_p mask
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(21)
logits=torch.randn(2,6,generator=g)
tk=R.top_k_mask(logits,3)
print('top_k dtype',tk.dtype)
print('logits',logits.tolist())
print('top_k',tk.tolist())
probs=torch.softmax(torch.randn(2,6,generator=g),-1)
tp=R.top_p_mask(probs,0.8)
print('probs',probs.tolist())
print('top_p',tp.tolist(),tp.dtype)
"
top_k dtype torch.bool logits [[0.10810236632823944, -0.4375609755516052, -0.7696864008903503, -0.19291412830352783, -0.3625756502151489, -2.8451075553894043], [1.4435173273086548, 0.49758392572402954, 0.6542161107063293, 0.07543526589870453, -1.076697826385498, 0.1269414722919464]] top_k [[True, False, False, True, True, False], [True, True, True, False, False, False]] probs [[0.43800121545791626, 0.03473725914955139, 0.04404238238930702, 0.02249596081674099, 0.07250185310840607, 0.38822129368782043], [0.3492460250854492, 0.07708215713500977, 0.39796486496925354, 0.08549511432647705, 0.025765402242541313, 0.06444645673036575]] top_p [[0.5301249623298645, 0.0, 0.0, 0.0, 0.0, 0.4698749780654907], [0.41941097378730774, 0.0, 0.4779176414012909, 0.10267142951488495, 0.0, 0.0]] torch.float32 [stdout] top_k dtype torch.bool logits [[0.10810236632823944, -0.4375609755516052, -0.7696864008903503, -0.19291412830352783, -0.3625756502151489, -2.8451075553894043], [1.4435173273086548, 0.49758392572402954, 0.6542161107063293, 0.07543526589870453, -1.076697826385498, 0.1269414722919464]] top_k [[True, False, False, True, True, False], [True, True, True, False, False, False]] probs [[0.43800121545791626, 0.03473725914955139, 0.04404238238930702, 0.02249596081674099, 0.07250185310840607, 0.38822129368782043], [0.3492460250854492, 0.07708215713500977, 0.39796486496925354, 0.08549511432647705, 0.025765402242541313, 0.06444645673036575]] top_p [[0.5301249623298645, 0.0, 0.0, 0.0, 0.0, 0.4698749780654907], [0.41941097378730774, 0.0, 0.4779176414012909, 0.10267142951488495, 0.0, 0.0]] torch.float32
Probe pad/nonzero/cummax/bucketize
python3 -c "
import torch, reforacle as R
lengths=torch.tensor([2,3,0,4])
print('pad_mask',R.pad_mask_from_lengths(lengths,4))
print('dtype',R.pad_mask_from_lengths(lengths,4).dtype)
mask=torch.tensor([[0,0,1,0,1.],[1,0,0,0,0],[0,0,0,0,0]])
print('first_nonzero',R.first_nonzero_index(mask))
x=torch.tensor([1.,3,2,5,4,0])
print('cummax',R.cumulative_max(x).tolist(), torch.cummax(x,0).values.tolist())
x2=torch.tensor([[1.,3,2],[5,1,9]])
print('cummax2',R.cumulative_max(x2).tolist())
rewards=torch.tensor([0.1,0.5,0.9,1.5,-0.3])
b=torch.tensor([0.0,0.5,1.0])
print('bucketize',R.bucketize_reward(rewards,b).tolist())
print('torch bucketize right F',torch.bucketize(rewards,b).tolist())
print('right T',torch.bucketize(rewards,b,right=True).tolist())
"
pad_mask tensor([[1., 1., 0., 0.],
[1., 1., 1., 0.],
[0., 0., 0., 0.],
[1., 1., 1., 1.]])
dtype torch.float32
first_nonzero tensor([ 2, 0, -1])
cummax [1.0, 1.0, 3.0, 3.0, 5.0, 5.0] [1.0, 3.0, 3.0, 5.0, 5.0, 5.0]
cummax2 [[1.0, 3.0, 2.0], [1.0, 3.0, 2.0]]
bucketize [1, 2, 2, 3, 0]
torch bucketize right F [1, 1, 2, 3, 0]
right T [1, 2, 2, 3, 0]
[stdout]
pad_mask tensor([[1., 1., 0., 0.],
[1., 1., 1., 0.],
[0., 0., 0., 0.],
[1., 1., 1., 1.]])
dtype torch.float32
first_nonzero tensor([ 2, 0, -1])
cummax [1.0, 1.0, 3.0, 3.0, 5.0, 5.0] [1.0, 3.0, 3.0, 5.0, 5.0, 5.0]
cummax2 [[1.0, 3.0, 2.0], [1.0, 3.0, 2.0]]
bucketize [1, 2, 2, 3, 0]
torch bucketize right F [1, 1, 2, 3, 0]
right T [1, 2, 2, 3, 0]Probe logprob_at_temperature/mode
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(22)
logits=torch.randn(2,4,6,generator=g)
labels=torch.randint(0,6,(2,4),generator=g)
T=0.7
lt=R.logprob_at_temperature(logits,labels,T)
print('shape',lt.shape)
print('ref',lt.tolist())
lsm=torch.log_softmax(logits/T,-1).gather(-1,labels.unsqueeze(-1)).squeeze(-1)
print('scaled',lsm.tolist())
# mode
labels2=torch.tensor([[1,2,2,3],[0,0,1,1]])
print('mode',R.mode_label(labels2).tolist(), torch.mode(labels2,dim=-1).values.tolist())
"
shape torch.Size([2, 4]) ref [[-1.9818044900894165, -4.526905536651611, -2.1092963218688965, -4.0334792137146], [-2.4389309883117676, -1.540858507156372, -2.466175079345703, -1.9439051151275635]] scaled [[-1.9818044900894165, -4.526905536651611, -2.1092963218688965, -4.0334792137146], [-2.4389309883117676, -1.540858507156372, -2.466175079345703, -1.9439051151275635]] mode [2, 1] [2, 0] [stdout] shape torch.Size([2, 4]) ref [[-1.9818044900894165, -4.526905536651611, -2.1092963218688965, -4.0334792137146], [-2.4389309883117676, -1.540858507156372, -2.466175079345703, -1.9439051151275635]] scaled [[-1.9818044900894165, -4.526905536651611, -2.1092963218688965, -4.0334792137146], [-2.4389309883117676, -1.540858507156372, -2.466175079345703, -1.9439051151275635]] mode [2, 1] [2, 0]
Probe mode tie-breaking
python3 -c " import torch, reforacle as R cases=[[0,0,1,1],[2,2,5,5],[3,3,3,1,1],[7,7,2,2,2],[1,2,3],[5,5,1,1,9,9]] for c in cases: t=torch.tensor([c]) print(c,'->',R.mode_label(t).item()) "
[0, 0, 1, 1] -> 1 [2, 2, 5, 5] -> 5 [3, 3, 3, 1, 1] -> 3 [7, 7, 2, 2, 2] -> 2 [1, 2, 3] -> 3 [5, 5, 1, 1, 9, 9] -> 9 [stdout] [0, 0, 1, 1] -> 1 [2, 2, 5, 5] -> 5 [3, 3, 3, 1, 1] -> 3 [7, 7, 2, 2, 2] -> 2 [1, 2, 3] -> 3 [5, 5, 1, 1, 9, 9] -> 9
Probe dpo_sequence_loss
python3 -c "
import torch, reforacle as R
import torch.nn.functional as F
g=torch.Generator().manual_seed(30)
V=6;B=3;Tk=4
pc=torch.randn(B,Tk,V,generator=g); pr=torch.randn(B,Tk,V,generator=g)
rc=torch.randn(B,Tk,V,generator=g); rr=torch.randn(B,Tk,V,generator=g)
cl=torch.randint(0,V,(B,Tk),generator=g); rl=torch.randint(0,V,(B,Tk),generator=g)
cm=(torch.rand(B,Tk,generator=g)>0.3).float(); rm=(torch.rand(B,Tk,generator=g)>0.3).float()
beta=0.1;ls=0.0
ref=R.dpo_sequence_loss(pc,pr,rc,rr,cl,rl,cm,rm,beta,ls)
print('ref',ref.item())
def seqlp(logits,labels,mask):
lp=torch.log_softmax(logits,-1).gather(-1,labels.unsqueeze(-1)).squeeze(-1)
return (lp*mask).sum(-1)
pcl=seqlp(pc,cl,cm); prl=seqlp(pr,rl,rm); rcl=seqlp(rc,cl,cm); rrl=seqlp(rr,rl,rm)
li=beta*((pcl-rcl)-(prl-rrl))
print('dpo from seq', (-F.logsigmoid(li)).mean().item())
"
ref 0.7570225596427917 dpo from seq 0.7570225596427917 [stdout] ref 0.7570225596427917 dpo from seq 0.7570225596427917
Probe grpo_objective baseline
python3 -c "
import torch, reforacle as R
import inspect
g=torch.Generator().manual_seed(31)
V=5;B=6;Tk=4
logits=torch.randn(B,Tk,V,generator=g)
old=torch.randn(B,Tk,V,generator=g)
ref_logits=torch.randn(B,Tk,V,generator=g)
labels=torch.randint(0,V,(B,Tk),generator=g)
cm=(torch.rand(B,Tk,generator=g)>0.3).float()
rewards=torch.randn(B,generator=g)
gs=3;beta=0.04;cl=0.2;ch=0.2
ref=R.grpo_objective(logits,old,ref_logits,labels,cm,rewards,gs,beta,cl,ch,True,'k3')
print('grpo ref',ref.item())
" 2>&1 | headgrpo ref 0.6902254223823547 [stdout] grpo ref 0.6902254223823547
Reconstruct grpo_objective
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(31)
V=5;B=6;Tk=4
logits=torch.randn(B,Tk,V,generator=g)
old=torch.randn(B,Tk,V,generator=g)
ref_logits=torch.randn(B,Tk,V,generator=g)
labels=torch.randint(0,V,(B,Tk),generator=g)
cm=(torch.rand(B,Tk,generator=g)>0.3).float()
rewards=torch.randn(B,generator=g)
gs=3;beta=0.04;cl=0.2;ch=0.2
def tlp(lg): return torch.log_softmax(lg,-1).gather(-1,labels.unsqueeze(-1)).squeeze(-1)
lp=tlp(logits); olp=tlp(old); rlp=tlp(ref_logits)
adv=R.grpo_advantages(rewards,gs,True) # [B]
A=adv.unsqueeze(-1)
ratio=torch.exp(lp-olp)
pg=-torch.min(ratio*A, torch.clamp(ratio,1-cl,1+ch)*A)
kl=torch.exp(rlp-lp)-1-(rlp-lp) # k3
tok=pg+beta*kl
# masked mean over all
print('all', ((tok*cm).sum()/cm.sum()).item())
# per-seq masked mean then mean
per=(tok*cm).sum(-1)/cm.sum(-1).clamp(min=1)
print('perseq mean', per.mean().item())
print('ref',R.grpo_objective(logits,old,ref_logits,labels,cm,rewards,gs,beta,cl,ch,True,'k3').item())
"
all 0.6902254223823547 perseq mean 0.3338797092437744 ref 0.6902254223823547 [stdout] all 0.6902254223823547 perseq mean 0.3338797092437744 ref 0.6902254223823547
Reconstruct ppo_objective
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(32)
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)
old_logp=torch.randn(T,generator=g)
nextv=torch.tensor(0.3)
gamma=0.99;lam=0.95;cl=0.2;ch=0.2;vfclip=0.2;vfcoef=0.5
ref=R.ppo_objective(rewards,values,old_values,logp,old_logp,nextv,gamma,lam,cl,ch,vfclip,vfcoef)
print('ref',ref.item())
# reconstruct
adv=R.gae(rewards,values,nextv,gamma,lam)
returns=adv+values
# whiten advantages? maybe
ratio=torch.exp(logp-old_logp)
pg=-torch.min(ratio*adv, torch.clamp(ratio,1-cl,1+ch)*adv).mean()
vc=old_values+torch.clamp(values-old_values,-vfclip,vfclip)
vl=0.5*torch.max((values-returns)**2,(vc-returns)**2).mean()
print('no whiten',(pg+vfcoef*vl).item())
# with whitened adv (no mask)
n=adv.numel();mean=adv.mean();var=((adv-mean)**2).mean()
advw=(adv-mean)*torch.rsqrt(var+1e-8)
pg2=-torch.min(ratio*advw, torch.clamp(ratio,1-cl,1+ch)*advw).mean()
print('whiten adv',(pg2+vfcoef*vl).item())
"
ref 0.8653490543365479 no whiten 0.583116888999939 whiten adv 0.8653490543365479 [stdout] ref 0.8653490543365479 no whiten 0.583116888999939 whiten adv 0.8653490543365479
Reconstruct rloo_objective
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(33)
V=5;B=6;Tk=4
logits=torch.randn(B,Tk,V,generator=g)
old=torch.randn(B,Tk,V,generator=g)
labels=torch.randint(0,V,(B,Tk),generator=g)
mask=(torch.rand(B,Tk,generator=g)>0.3).float()
rewards=torch.randn(B,generator=g)
gs=3;cl=0.2;ch=0.2
ref=R.rloo_objective(logits,old,labels,mask,rewards,gs,cl,ch)
print('rloo ref',ref.item())
def tlp(lg): return torch.log_softmax(lg,-1).gather(-1,labels.unsqueeze(-1)).squeeze(-1)
lp=tlp(logits);olp=tlp(old)
adv=R.rloo_advantages(rewards,gs).unsqueeze(-1)
ratio=torch.exp(lp-olp)
pg=-torch.min(ratio*adv,torch.clamp(ratio,1-cl,1+ch)*adv)
print('masked mean',((pg*mask).sum()/mask.sum()).item())
"
rloo ref 0.8057166337966919 masked mean 0.8057166337966919 [stdout] rloo ref 0.8057166337966919 masked mean 0.8057166337966919
Reconstruct reverse_kl_objective
python3 -c "
import torch, reforacle as R
g=torch.Generator().manual_seed(34)
V=5;B=3;Tk=4
logits=torch.randn(B,Tk,V,generator=g)
ref_logits=torch.randn(B,Tk,V,generator=g)
labels=torch.randint(0,V,(B,Tk),generator=g)
mask=(torch.rand(B,Tk,generator=g)>0.3).float()
advantages=torch.randn(B,generator=g)
beta=0.1
ref=R.reverse_kl_objective(logits,ref_logits,labels,mask,advantages,beta)
print('ref',ref.item())
def tlp(lg): return torch.log_softmax(lg,-1).gather(-1,labels.unsqueeze(-1)).squeeze(-1)
lp=tlp(logits);rlp=tlp(ref_logits)
A=advantages.unsqueeze(-1)
kl=torch.exp(rlp-lp)-1-(rlp-lp)
# pg = -advantage * logp ? policy gradient
pg=-A*lp
tok=pg+beta*kl
print('masked mean pg+beta kl',((tok*mask).sum()/mask.sum()).item())
tok2=-A*lp
print('just pg',((tok2*mask).sum()/mask.sum()).item())
tok3=-(A*lp)+beta*kl
print('same',((tok3*mask).sum()/mask.sum()).item())
"
ref -1.5088324546813965 masked mean pg+beta kl -1.5088324546813965 just pg -1.5464565753936768 same -1.5088324546813965 [stdout] ref -1.5088324546813965 masked mean pg+beta kl -1.5088324546813965 just pg -1.5464565753936768 same -1.5088324546813965
from __future__ import annotations
import math
import torch
def masked_mean(values, mask, axis=None):
mask = mask.to(values.dtype)
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=1.0)
def masked_sum(values, mask, axis=None):
mask = mask.to(values.dtype)
if axis is None:
return (values * mask).sum()
return (values * mask).sum(axis)
def logsumexp(x, axis):
return torch.logsumexp(x, axis)
def log_softmax(x, axis):
return torch.log_softmax(x, axis)
def token_logprobs(logits, labels):
lsm = torch.log_softmax(logits, -1)
return lsm.gather(-1, labels.unsqueeze(-1)).squeeze(-1)
def selective_logprobs(logits, labels, mask):
lp = token_logprobs(logits, labels)
mask = mask.to(lp.dtype)
return (lp * mask).sum(-1)
def sequence_logprob(logits, labels, mask, length_normalize):
lp = token_logprobs(logits, labels)
mask = mask.to(lp.dtype)
seq = (lp * mask).sum(-1)
if length_normalize:
seq = seq / mask.sum(-1).clamp(min=1.0)
return seq
def entropy(logits, mask):
lsm = torch.log_softmax(logits, -1)
p = lsm.exp()
ent = -(p * lsm).sum(-1)
return masked_mean(ent, mask)
def dpo_loss(pc, pr, rc, rr, beta, label_smoothing):
logits = beta * ((pc - rc) - (pr - rr))
loss = (-torch.nn.functional.logsigmoid(logits) * (1 - label_smoothing)
- torch.nn.functional.logsigmoid(-logits) * label_smoothing)
return loss.mean()
def ipo_loss(pc, pr, rc, rr, beta):
h = (pc - pr) - (rc - rr)
return ((h - 1.0 / (2.0 * beta)) ** 2).mean()
def grpo_advantages(rewards, group_size, scale_by_std):
shape = rewards.shape
grouped = rewards.reshape(-1, group_size)
mean = grouped.mean(dim=-1, keepdim=True)
adv = grouped - mean
if scale_by_std:
std = grouped.std(dim=-1, unbiased=False, keepdim=True)
adv = adv / (std + 1e-8)
return adv.reshape(shape)
def gae(rewards, values, next_value, gamma, lam):
adv = torch.zeros_like(rewards)
T = rewards.shape[0]
last = torch.zeros_like(next_value)
for t in reversed(range(T)):
nv = values[t + 1] if t + 1 < T else next_value
delta = rewards[t] + gamma * nv - 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
if estimator == "k3":
return torch.exp(-d) - 1.0 + d
raise ValueError(estimator)
def clipped_pg_loss(logp, old_logp, advantages, mask, clip_low, clip_high):
ratio = torch.exp(logp - old_logp)
unclipped = ratio * advantages
clipped = torch.clamp(ratio, 1.0 - clip_low, 1.0 + clip_high) * advantages
loss = -torch.min(unclipped, clipped)
return masked_mean(loss, mask)
def value_loss(values, old_values, returns, clip):
vclipped = old_values + torch.clamp(values - old_values, -clip, clip)
l1 = (values - returns) ** 2
l2 = (vclipped - returns) ** 2
return 0.5 * torch.max(l1, l2).mean()
def _whiten_stats(values, mask):
mask = mask.to(values.dtype)
n = mask.sum()
mean = (values * mask).sum() / n.clamp(min=1.0)
var = ((values - mean) ** 2 * mask).sum() / n.clamp(min=1.0)
return mean, var
def whiten(values, mask, shift_mean):
mean, var = _whiten_stats(values, mask)
whitened = (values - mean) * torch.rsqrt(var + 1e-8)
if not shift_mean:
whitened = whitened + mean
return whitened
def masked_whiten(values, mask, shift_mean):
mean, var = _whiten_stats(values, mask)
whitened = (values - mean) * torch.rsqrt(var + 1e-8) * mask.to(values.dtype)
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 = selective_logprobs(pc_logits, chosen_labels, chosen_mask)
pr = selective_logprobs(pr_logits, rejected_labels, rejected_mask)
rc = selective_logprobs(rc_logits, chosen_labels, chosen_mask)
rr = selective_logprobs(rr_logits, rejected_labels, rejected_mask)
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):
lp = token_logprobs(logits, labels)
olp = token_logprobs(old_logits, labels)
rlp = token_logprobs(ref_logits, labels)
adv = grpo_advantages(rewards, group_size, scale_by_std).unsqueeze(-1)
ratio = torch.exp(lp - olp)
pg = -torch.min(ratio * adv, torch.clamp(ratio, 1.0 - clip_low, 1.0 + clip_high) * adv)
kl = kl_penalty(lp, rlp, 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
adv = whiten(adv, torch.ones_like(adv), True)
ratio = torch.exp(logp - old_logp)
pg = -torch.min(ratio * adv, torch.clamp(ratio, 1.0 - clip_low, 1.0 + clip_high) * adv).mean()
vl = value_loss(values, old_values, returns, vf_clip)
return pg + vf_coef * vl
def rloo_advantages(rewards, group_size):
shape = rewards.shape
grouped = rewards.reshape(-1, group_size)
total = grouped.sum(dim=-1, keepdim=True)
baseline = (total - grouped) / (group_size - 1)
return (grouped - baseline).reshape(shape)
def reverse_kl(logp, ref_logp):
d = logp - ref_logp
return torch.exp(-d) - 1.0 + d
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):
out = torch.zeros_like(rewards)
T = rewards.shape[0]
acc = torch.zeros_like(rewards[0]) if T > 0 else None
for t in reversed(range(T)):
acc = rewards[t] + gamma * acc
out[t] = acc
return out
def normalize(x, eps):
mean = x.mean()
std = x.std()
return (x - mean) / (std + eps)
def top_p_mask(probs, p):
sorted_probs, sorted_idx = torch.sort(probs, dim=-1, descending=True)
cum = torch.cumsum(sorted_probs, dim=-1)
remove = (cum - sorted_probs) > p
sorted_probs = sorted_probs.masked_fill(remove, 0.0)
out = torch.zeros_like(probs)
out.scatter_(-1, sorted_idx, sorted_probs)
out = out / out.sum(dim=-1, keepdim=True)
return out
def smoothed_nll(logits, labels, smoothing):
lsm = torch.log_softmax(logits, -1)
nll = -lsm.gather(-1, labels.unsqueeze(-1)).squeeze(-1)
smooth = -lsm.mean(-1)
loss = (1 - 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):
lp = token_logprobs(logits, labels)
olp = token_logprobs(old_logits, labels)
adv = rloo_advantages(rewards, group_size).unsqueeze(-1)
ratio = torch.exp(lp - olp)
pg = -torch.min(ratio * adv, torch.clamp(ratio, 1.0 - clip_low, 1.0 + clip_high) * adv)
return masked_mean(pg, mask)
def reverse_kl_objective(logits, ref_logits, labels, mask, advantages, beta):
lp = token_logprobs(logits, labels)
rlp = token_logprobs(ref_logits, labels)
adv = advantages.unsqueeze(-1)
kl = reverse_kl(lp, rlp)
tok = -adv * lp + beta * kl
return masked_mean(tok, mask)
def cross_entropy(logits, labels, ignore_index):
V = logits.shape[-1]
return torch.nn.functional.cross_entropy(
logits.reshape(-1, V), labels.reshape(-1), ignore_index=ignore_index)
def top_k_mask(logits, k):
topk = torch.topk(logits, k, dim=-1).values
threshold = topk[..., -1:].expand_as(logits)
return logits >= threshold
def group_mean_baseline(rewards, group_size):
return grpo_advantages(rewards, group_size, False)
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
fwd = torch.exp(d) - 1.0 - d
rev = torch.exp(-d) - 1.0 + d
return 0.5 * (fwd + rev)
def huber_value_loss(values, returns, delta):
return torch.nn.functional.huber_loss(values, returns, delta=delta)
def normalized_entropy(logits, mask):
V = logits.shape[-1]
return entropy(logits, mask) / math.log(V)
def clip_fraction(logp, old_logp, clip):
ratio = torch.exp(logp - old_logp)
return ((ratio - 1.0).abs() > clip).float().mean()
def logprob_at_temperature(logits, labels, temperature):
lsm = torch.log_softmax(logits / temperature, -1)
return lsm.gather(-1, labels.unsqueeze(-1)).squeeze(-1)
def advantage_mean_std(advantages, mask):
mean, var = _whiten_stats(advantages, mask)
std = torch.sqrt(var)
return torch.stack([mean, std])
def argmax_tokens(logits):
return logits.argmax(-1)
def mode_label(labels):
out_shape = labels.shape[:-1]
flat = labels.reshape(-1, labels.shape[-1])
result = torch.empty(flat.shape[0], dtype=labels.dtype, device=labels.device)
for i in range(flat.shape[0]):
vals, counts = torch.unique(flat[i], return_counts=True)
maxc = counts.max()
cand = vals[counts == maxc]
result[i] = cand.max()
return result.reshape(out_shape)
def median_reward(rewards):
if rewards.dim() == 0:
return rewards
sr = rewards.sort(dim=-1).values
n = sr.shape[0]
pos = 0.5 * (n - 1)
lo = math.floor(pos)
hi = math.ceil(pos)
frac = pos - lo
return sr[lo] * (1 - frac) + sr[hi] * frac
def masked_max(values, mask):
filled = torch.where(mask > 0, values, torch.full_like(values, float("-inf")))
return filled.max()
def quantile_lower(x, q):
return torch.quantile(x, q, interpolation="lower")
def pad_mask_from_lengths(lengths, max_len):
ar = torch.arange(max_len, device=lengths.device)
return (ar < lengths.unsqueeze(-1)).to(torch.float32)
def first_nonzero_index(mask):
nonzero = mask != 0
any_nz = nonzero.any(dim=-1)
idx = torch.argmax(nonzero.to(torch.long), dim=-1)
return torch.where(any_nz, idx, torch.full_like(idx, -1))
def cumulative_max(x):
c = torch.cummax(x, dim=0).values
return torch.cat([x[:1], c[:-1]], dim=0)
def bucketize_reward(rewards, boundaries):
return torch.bucketize(rewards, boundaries, right=True)
import torch, math
import reforacle as R
import objectives as O
torch.manual_seed(0)
fails = {}
def cmp(name, args_list, atol=1e-6, rtol=1e-5):
ok = True
for args in args_list:
try:
r = getattr(R, name)(*[a.clone() if torch.is_tensor(a) else a for a in args])
except Exception as e:
r = ('ERR', type(e).__name__)
try:
m = getattr(O, name)(*[a.clone() if torch.is_tensor(a) else a for a in args])
except Exception as e:
m = ('ERR', type(e).__name__)
if isinstance(r, tuple) or isinstance(m, tuple):
if r != m:
ok = False; print(f" {name} mismatch exc {r} vs {m}")
continue
if not torch.is_tensor(r): r = torch.tensor(r)
if not torch.is_tensor(m): m = torch.tensor(m)
if r.shape != m.shape:
ok = False; print(f" {name} SHAPE {r.shape} vs {m.shape}"); continue
if r.dtype != m.dtype:
print(f" {name} DTYPE {r.dtype} vs {m.dtype}")
if not torch.allclose(r.float(), m.float(), atol=atol, rtol=rtol, equal_nan=True):
ok = False
diff = (r.float()-m.float()).abs().max().item()
print(f" {name} NUM maxdiff={diff}")
print(("PASS" if ok else "FAIL"), name)
if not ok: fails[name] = True
def rn(*s): return torch.randn(*s)
def rm(*s): return (torch.rand(*s) > 0.4).float()
def ri(hi,*s): return torch.randint(0,hi,s)
# primitives
cmp('masked_mean', [(rn(8), rm(8)), (rn(3,4), rm(3,4), 1), (rn(3,4), rm(3,4), 0), (rn(3,4), rm(3,4), None), (rn(5,6,7), rm(5,6,7), 2)])
cmp('masked_sum', [(rn(8), rm(8)), (rn(3,4), rm(3,4), 1), (rn(3,4), rm(3,4), 0), (rn(3,4), rm(3,4), None)])
cmp('logsumexp', [(rn(3,5),1),(rn(3,5),0),(rn(2,3,4),2),(rn(2,3,4),-1)])
cmp('log_softmax', [(rn(3,5),1),(rn(2,3,4),-1),(rn(2,3,4),0)])
cmp('token_logprobs', [(rn(2,4,6), ri(6,2,4)), (rn(5,7), ri(7,5))])
cmp('selective_logprobs', [(rn(2,4,6), ri(6,2,4), rm(2,4))])
cmp('sequence_logprob', [(rn(2,4,6), ri(6,2,4), rm(2,4), False), (rn(2,4,6), ri(6,2,4), rm(2,4), True)])
cmp('entropy', [(rn(2,4,6), rm(2,4)), (rn(3,5,7), rm(3,5))])
cmp('dpo_loss', [(rn(4),rn(4),rn(4),rn(4),0.1,0.0),(rn(4),rn(4),rn(4),rn(4),0.5,0.1)])
cmp('ipo_loss', [(rn(4),rn(4),rn(4),rn(4),0.1),(rn(5),rn(5),rn(5),rn(5),0.3)])
cmp('grpo_advantages', [(rn(6),3,False),(rn(6),3,True),(rn(12),4,True),(rn(12),4,False)])
cmp('gae', [(rn(5),rn(5),torch.tensor(0.3),0.99,0.95),(rn(4,2),rn(4,2),rn(2),0.9,0.9)])
cmp('kl_penalty', [(rn(5),rn(5),'k1'),(rn(5),rn(5),'k2'),(rn(5),rn(5),'k3')])
cmp('clipped_pg_loss', [(rn(2,4),rn(2,4),rn(2,4),rm(2,4),0.2,0.2)])
cmp('value_loss', [(rn(2,4),rn(2,4),rn(2,4),0.2)])
cmp('whiten', [(rn(8),rm(8),True),(rn(8),rm(8),False)])
cmp('masked_whiten', [(rn(8),rm(8),True),(rn(8),rm(8),False)])
V,B,Tk=6,6,4
cmp('dpo_sequence_loss', [(rn(3,Tk,V),rn(3,Tk,V),rn(3,Tk,V),rn(3,Tk,V),ri(V,3,Tk),ri(V,3,Tk),rm(3,Tk),rm(3,Tk),0.1,0.0)], atol=1e-5, rtol=1e-4)
cmp('grpo_objective', [(rn(B,Tk,V),rn(B,Tk,V),rn(B,Tk,V),ri(V,B,Tk),rm(B,Tk),rn(B),3,0.04,0.2,0.2,True,'k3'),
(rn(B,Tk,V),rn(B,Tk,V),rn(B,Tk,V),ri(V,B,Tk),rm(B,Tk),rn(B),2,0.1,0.2,0.3,False,'k1')], atol=1e-5, rtol=1e-4)
cmp('ppo_objective', [(rn(5),rn(5),rn(5),rn(5),rn(5),torch.tensor(0.3),0.99,0.95,0.2,0.2,0.2,0.5)], atol=1e-5, rtol=1e-4)
cmp('rloo_advantages', [(rn(6),3),(rn(12),4)])
cmp('reverse_kl', [(rn(5),rn(5))])
cmp('importance_ratio', [(rn(5),rn(5),None),(rn(5),rn(5),0.2)])
cmp('discounted_returns', [(rn(5),0.99),(rn(4,2),0.9)])
cmp('normalize', [(rn(8),1e-8),(rn(10),1e-6)])
cmp('top_p_mask', [(torch.softmax(rn(2,6),-1),0.8),(torch.softmax(rn(3,8),-1),0.5),(torch.softmax(rn(4,5),-1),0.9)], atol=1e-5, rtol=1e-4)
cmp('smoothed_nll', [(rn(2,4,6),ri(6,2,4),0.1),(rn(3,5),ri(5,3),0.0)])
cmp('bradley_terry_logit', [(torch.tensor(1.0),torch.tensor(0.3),0.5)])
cmp('rloo_objective', [(rn(B,Tk,V),rn(B,Tk,V),ri(V,B,Tk),rm(B,Tk),rn(B),3,0.2,0.2)], atol=1e-5, rtol=1e-4)
cmp('reverse_kl_objective', [(rn(3,Tk,V),rn(3,Tk,V),ri(V,3,Tk),rm(3,Tk),rn(3),0.1)], atol=1e-5, rtol=1e-4)
cmp('cross_entropy', [(rn(3,5,7),ri(7,3,5),-100),(rn(4,6),ri(6,4),-100)])
cmp('top_k_mask', [(rn(2,6),3),(rn(3,8),2),(rn(4,5),1)])
cmp('group_mean_baseline', [(rn(6),3),(rn(12),4)])
cmp('lambda_returns', [(rn(5),rn(5),torch.tensor(0.3),0.99,0.95),(rn(4,2),rn(4,2),rn(2),0.9,0.9)])
cmp('symmetric_kl', [(rn(5),rn(5))])
cmp('huber_value_loss', [(rn(2,4),rn(2,4),1.0),(rn(3,5),rn(3,5),0.5)])
cmp('normalized_entropy', [(rn(2,4,6),rm(2,4))])
cmp('clip_fraction', [(rn(20)*0.2,rn(20)*0.2,0.2),(rn(10),rn(10),0.3)])
cmp('logprob_at_temperature', [(rn(2,4,6),ri(6,2,4),0.7),(rn(3,5),ri(5,3),1.5)])
cmp('advantage_mean_std', [(rn(2,4),rm(2,4)),(rn(10),rm(10))])
cmp('argmax_tokens', [(rn(2,4,6),),(rn(3,5),)])
cmp('mode_label', [(ri(5,2,4),),(ri(3,6),),(torch.tensor([[0,0,1,1],[2,2,2,3]]),)])
cmp('median_reward', [(rn(5),),(rn(3,4),),(rn(6,3),),(rn(8),)])
cmp('masked_max', [(rn(2,4),rm(2,4)),(rn(10),rm(10))])
cmp('quantile_lower', [(rn(10),0.25),(rn(20),0.5),(rn(15),0.75)])
cmp('pad_mask_from_lengths', [(torch.tensor([2,3,0,4]),4),(torch.tensor([1,5,2]),5)])
cmp('first_nonzero_index', [(torch.tensor([[0,0,1,0,1],[1,0,0,0,0],[0,0,0,0,0]]).float(),),(rm(4,6),)])
cmp('cumulative_max', [(rn(6),),(rn(4,3),)])
cmp('bucketize_reward', [(rn(6),torch.tensor([0.0,0.5,1.0])),(rn(8),torch.linspace(-1,1,4))])
print()
print("TOTAL FAILS:", list(fails.keys()) if fails else "NONE")
Run comprehensive test
python3 test_all.py 2>&1
PASS masked_mean PASS masked_sum PASS logsumexp PASS log_softmax PASS token_logprobs PASS selective_logprobs PASS sequence_logprob PASS entropy PASS dpo_loss PASS ipo_loss PASS grpo_advantages PASS gae PASS kl_penalty PASS clipped_pg_loss PASS value_loss PASS whiten PASS masked_whiten 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 PASS normalized_entropy PASS clip_fraction 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 TOTAL FAILS: NONE [stdout] PASS masked_mean PASS masked_sum PASS logsumexp PASS log_softmax PASS token_logprobs PASS selective_logprobs PASS sequence_logprob PASS entropy PASS dpo_loss PASS ipo_loss PASS grpo_advantages PASS gae PASS kl_penalty PASS clipped_pg_loss PASS value_loss PASS whiten PASS masked_whiten 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 PASS normalized_entropy PASS clip_fraction 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 TOTAL FAILS: NONE
Edge cases: ties, empty masks
python3 -c "
import torch, reforacle as R, objectives as O
# top_k ties
logits=torch.tensor([[1.0,1.0,1.0,0.0,2.0],[3.,3.,1.,1.,1.]])
print('topk ref',R.top_k_mask(logits,2).tolist())
print('topk mine',O.top_k_mask(logits,2).tolist())
# masked_mean fully masked row along axis
V=torch.randn(3,4); M=(torch.rand(3,4)>0.4).float(); M[1]=0
print('mm axis1 ref',R.masked_mean(V,M,1).tolist())
print('mm axis1 mine',O.masked_mean(V,M,1).tolist())
print('ms axis1 match',torch.allclose(R.masked_sum(V,M,1),O.masked_sum(V,M,1)))
# masked_max all zero mask
v=torch.randn(5); m=torch.zeros(5)
print('mmax empty ref',R.masked_max(v,m).item(),'mine',O.masked_max(v,m).item())
# first_nonzero ties within
print('fnz ref',R.first_nonzero_index(torch.tensor([[1,1,0.],[0,0,0]])).tolist())
print('fnz mine',O.first_nonzero_index(torch.tensor([[1,1,0.],[0,0,0]])).tolist())
"
topk ref [[True, True, True, False, True], [True, True, False, False, False]] topk mine [[True, True, True, False, True], [True, True, False, False, False]] mm axis1 ref [-0.043334562331438065, 0.0, -1.1354539394378662] mm axis1 mine [-0.043334562331438065, 0.0, -1.1354539394378662] ms axis1 match True mmax empty ref 0.0 mine -inf fnz ref [0, -1] fnz mine [0, -1] [stdout] topk ref [[True, True, True, False, True], [True, True, False, False, False]] topk mine [[True, True, True, False, True], [True, True, False, False, False]] mm axis1 ref [-0.043334562331438065, 0.0, -1.1354539394378662] mm axis1 mine [-0.043334562331438065, 0.0, -1.1354539394378662] ms axis1 match True mmax empty ref 0.0 mine -inf fnz ref [0, -1] fnz mine [0, -1]
Probe masked_max empty/negatives
python3 -c "
import torch, reforacle as R
# all negative values, mask picks subset
v=torch.tensor([-5.,-2.,-8.,-1.,-9.])
m=torch.tensor([1.,0.,1.,0.,1.])
print('ref',R.masked_max(v,m).item()) # max of -5,-8,-9 = -5
m2=torch.tensor([0.,0.,1.,0.,0.])
print('ref single -8',R.masked_max(v,m2).item())
print('ref empty',R.masked_max(v,torch.zeros(5)).item())
# 2d?
V=torch.randn(2,3); M=torch.zeros(2,3)
print('ref empty 2d',R.masked_max(V,M))
"
ref -5.0 ref single -8 -8.0 ref empty 0.0 ref empty 2d tensor(0.) [stdout] ref -5.0 ref single -8 -8.0 ref empty 0.0 ref empty 2d tensor(0.)
/app/objectives.py
filled = torch.where(mask > 0, values, torch.full_like(values, float("-inf")))
return filled.max() filled = torch.where(mask > 0, values, torch.full_like(values, float("-inf")))
return torch.nan_to_num(filled.max(), neginf=0.0)Verify masked_max fix
python3 -c "
import torch, reforacle as R, objectives as O
for m in [torch.zeros(5), torch.tensor([1.,0,1,0,1]), torch.ones(5)]:
v=torch.randn(5)
print(torch.allclose(R.masked_max(v,m),O.masked_max(v,m)), R.masked_max(v,m).item(), O.masked_max(v,m).item())
V=torch.randn(2,3)
print('2d empty',R.masked_max(V,torch.zeros(2,3)).item(),O.masked_max(V,torch.zeros(2,3)).item())
"
True 0.0 0.0 True 0.5121377110481262 0.5121377110481262 True 0.9216684699058533 0.9216684699058533 2d empty 0.0 0.0 [stdout] True 0.0 0.0 True 0.5121377110481262 0.5121377110481262 True 0.9216684699058533 0.9216684699058533 2d empty 0.0 0.0
Edge probes batch
python3 -c "
import torch, reforacle as R, objectives as O
def ok(a,b):
a=a if torch.is_tensor(a) else torch.tensor(a); b=b if torch.is_tensor(b) else torch.tensor(b)
return a.shape==b.shape and torch.allclose(a.float(),b.float(),atol=1e-6,rtol=1e-5,equal_nan=True)
# seq logprob empty mask row
lg=torch.randn(2,4,6); lb=torch.randint(0,6,(2,4)); m=(torch.rand(2,4)>0.4).float(); m[0]=0
print('seq_ln empty',ok(R.sequence_logprob(lg,lb,m,True),O.sequence_logprob(lg,lb,m,True)))
# top_p extremes
for p in [0.0,0.01,1.0,0.999]:
pr=torch.softmax(torch.randn(3,7),-1)
print('top_p',p,ok(R.top_p_mask(pr,p),O.top_p_mask(pr,p)))
# single dominant
pr=torch.tensor([[0.9,0.05,0.03,0.02]])
print('top_p dom 0.5',ok(R.top_p_mask(pr,0.5),O.top_p_mask(pr,0.5)))
# mode 3D
lbl=torch.randint(0,4,(2,3,5))
print('mode3d',ok(R.mode_label(lbl),O.mode_label(lbl)))
# first_nonzero 1d
print('fnz1d',ok(R.first_nonzero_index(torch.tensor([0,0,1,0.])),O.first_nonzero_index(torch.tensor([0,0,1,0.]))))
print('fnz1d val',R.first_nonzero_index(torch.tensor([0,0,1,0.])), O.first_nonzero_index(torch.tensor([0,0,1,0.])))
# gae 3d, discounted 3d
r=torch.randn(4,2,3);v=torch.randn(4,2,3);nv=torch.randn(2,3)
print('gae3d',ok(R.gae(r,v,nv,0.9,0.9),O.gae(r,v,nv,0.9,0.9)))
print('disc3d',ok(R.discounted_returns(r,0.9),O.discounted_returns(r,0.9)))
# cross_entropy 2d simple (N,C)
print('ce2d',ok(R.cross_entropy(torch.randn(5,7),torch.randint(0,7,(5,)),-100),O.cross_entropy(torch.randn(5,7),torch.randint(0,7,(5,)),-100)))
"
seq_ln empty True top_p 0.0 True top_p 0.01 True top_p 1.0 True top_p 0.999 True top_p dom 0.5 True mode3d True fnz1d True fnz1d val tensor(2) tensor(2) gae3d True disc3d True ce2d False [stdout] seq_ln empty True top_p 0.0 True top_p 0.01 True top_p 1.0 True top_p 0.999 True top_p dom 0.5 True mode3d True fnz1d True fnz1d val tensor(2) tensor(2) gae3d True disc3d True ce2d False
Debug cross_entropy 2d
python3 -c "
import torch, reforacle as R, objectives as O
torch.manual_seed(1)
lg=torch.randn(5,7); lb=torch.randint(0,7,(5,))
print('ref',R.cross_entropy(lg,lb,-100))
print('mine',O.cross_entropy(lg,lb,-100))
"
ref tensor(2.1003) mine tensor(2.1003) [stdout] ref tensor(2.1003) mine tensor(2.1003)
Final full test run
python3 test_all.py 2>&1 | tail -5; echo "---timer---"; bash /app/timer.sh
PASS first_nonzero_index PASS cumulative_max PASS bucketize_reward TOTAL FAILS: NONE ---timer--- Remaining time (hours:minutes): 1:46 [stdout] PASS first_nonzero_index PASS cumulative_max PASS bucketize_reward TOTAL FAILS: NONE ---timer--- Remaining time (hours:minutes): 1:46
Verify standalone, no reforacle dependency
cd /tmp && grep -n "reforacle\|import" /app/objectives.py; echo "--- clean-sandbox import test ---"; python3 -c "
import sys, importlib
# simulate absence of reforacle
sys.modules['reforacle']=None
import importlib.util
spec=importlib.util.spec_from_file_location('obj','/app/objectives.py')
m=importlib.util.module_from_spec(spec); spec.loader.exec_module(m)
import torch
print('loads fine; masked_mean ->', m.masked_mean(torch.tensor([1.,2,3]),torch.tensor([1.,0,1])).item())
print('num functions:', len([x for x in dir(m) if not x.startswith('_') and callable(getattr(m,x))]))
"1:from __future__ import annotations 3:import math 5:import torch 193:def importance_ratio(logp, old_logp, clip): --- clean-sandbox import test --- loads fine; masked_mean -> 2.0 num functions: 49 Shell cwd was reset to /app [stdout] 1:from __future__ import annotations 3:import math 5:import torch 193:def importance_ratio(logp, old_logp, clip): --- clean-sandbox import test --- loads fine; masked_mean -> 2.0 num functions: 49 [stderr] Shell cwd was reset to /app
No file changes, the agent left the workspace untouched.
=== SyncVals verifier (post-training objectives, sealed oracle) === FUNCTIONS PASSED: 44/49 (partial_score=0.898) 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_878c713faaba4a87. Re-running the agent live requires EVAL_PLATFORM_ENABLE_OAUTH_SMOKE=1 and is non-deterministic.
Trial trial_878c713faaba4a87 · verifier authoritative; classifier explanatory.