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']
import torch
import reforacle as R
g = torch.Generator().manual_seed(1)
def rn(*s): return torch.randn(*s, generator=g)
# masked_mean axis behavior
v = rn(3,4); m = (torch.rand(3,4,generator=g)>0.4).float()
print("mm none", R.masked_mean(v,m).item())
print("mm ax0", R.masked_mean(v,m,axis=0).tolist())
print("mm ax1", R.masked_mean(v,m,axis=1).tolist())
print("mm empty", R.masked_mean(rn(4), torch.zeros(4)).item())
print("ms none", R.masked_sum(v,m).item())
print("ms ax1", R.masked_sum(v,m,axis=1).tolist())
x = rn(3,5)
print("lse", R.logsumexp(x,1).tolist())
print("lsm shape", R.log_softmax(x,1).shape, R.log_softmax(x,1)[0].tolist())
# token_logprobs
logits = rn(2,3,5)
labels = torch.randint(0,5,(2,3),generator=g)
tl = R.token_logprobs(logits, labels)
print("token_logprobs shape", tl.shape, tl.tolist())
msk = (torch.rand(2,3,generator=g)>0.3).float()
sl = R.selective_logprobs(logits, labels, msk)
print("selective type", type(sl), sl.shape if hasattr(sl,'shape') else sl)
print("selective", sl.tolist() if hasattr(sl,'tolist') else sl)
sq = R.sequence_logprob(logits, labels, msk, False)
print("seq_lp nolen", sq.tolist() if hasattr(sq,'tolist') else sq)
sq2 = R.sequence_logprob(logits, labels, msk, True)
print("seq_lp len", sq2.tolist() if hasattr(sq2,'tolist') else sq2)
en = R.entropy(logits, msk)
print("entropy", type(en), en.tolist() if hasattr(en,'tolist') else en)
Run probe 1
python3 p1.py
mm none -0.1978435516357422 mm ax0 [-1.0276086330413818, 0.05039694160223007, -0.7305456399917603, 0.5015006065368652] mm ax1 [0.31663957238197327, -0.4357382357120514, -1.0276086330413818] mm empty 0.0 ms none -1.3849048614501953 ms ax1 [0.9499186873435974, -1.3072147369384766, -1.0276086330413818] lse [2.4159557819366455, 2.805713653564453, 1.8280694484710693] lsm shape torch.Size([3, 5]) [-0.8172554969787598, -3.692962646484375, -2.090454339981079, -2.8951005935668945, -1.03694748878479] token_logprobs shape torch.Size([2, 3]) [[-2.573004961013794, -1.458168387413025, -3.7930383682250977], [-1.3228989839553833, -2.2500669956207275, -3.3174071311950684]] selective type <class 'torch.Tensor'> torch.Size([2]) selective [-4.031173229217529, -6.890373229980469] seq_lp nolen [-4.031173229217529, -6.890373229980469] seq_lp len [-2.0155866146087646, -2.2967910766601562] entropy <class 'torch.Tensor'> 1.3083505630493164 [stdout] mm none -0.1978435516357422 mm ax0 [-1.0276086330413818, 0.05039694160223007, -0.7305456399917603, 0.5015006065368652] mm ax1 [0.31663957238197327, -0.4357382357120514, -1.0276086330413818] mm empty 0.0 ms none -1.3849048614501953 ms ax1 [0.9499186873435974, -1.3072147369384766, -1.0276086330413818] lse [2.4159557819366455, 2.805713653564453, 1.8280694484710693] lsm shape torch.Size([3, 5]) [-0.8172554969787598, -3.692962646484375, -2.090454339981079, -2.8951005935668945, -1.03694748878479] token_logprobs shape torch.Size([2, 3]) [[-2.573004961013794, -1.458168387413025, -3.7930383682250977], [-1.3228989839553833, -2.2500669956207275, -3.3174071311950684]] selective type <class 'torch.Tensor'> torch.Size([2]) selective [-4.031173229217529, -6.890373229980469] seq_lp nolen [-4.031173229217529, -6.890373229980469] seq_lp len [-2.0155866146087646, -2.2967910766601562] entropy <class 'torch.Tensor'> 1.3083505630493164
import torch
import reforacle as R
import torch.nn.functional as F
g = torch.Generator().manual_seed(2)
def rn(*s): return torch.randn(*s, generator=g)
# entropy check: is it masked_mean of per-token entropy?
logits = rn(2,3,5)
msk = (torch.rand(2,3,generator=g)>0.3).float()
lp = torch.log_softmax(logits,-1); p=lp.exp()
ent_tok = -(p*lp).sum(-1)
print("my masked_mean ent", (ent_tok*msk).sum()/msk.sum())
print("R entropy", R.entropy(logits,msk).item())
print("R normalized_entropy", R.normalized_entropy(logits,msk).item())
print("norm guess ent/log(V)", ((ent_tok*msk).sum()/msk.sum()/torch.log(torch.tensor(5.0))).item())
# dpo
pc,pr,rc,rr = rn(4),rn(4),rn(4),rn(4)
print("dpo b0.1 ls0", R.dpo_loss(pc,pr,rc,rr,0.1,0.0).item())
print("dpo b0.5 ls0.1", R.dpo_loss(pc,pr,rc,rr,0.5,0.1).item())
# my dpo
beta=0.1;ls=0.0
li = beta*((pc-rc)-(pr-rr))
mine = (-F.logsigmoid(li)*(1-ls) - F.logsigmoid(-li)*ls).mean()
print("my dpo b0.1", mine.item())
beta=0.5;ls=0.1
li = beta*((pc-rc)-(pr-rr))
mine = (-F.logsigmoid(li)*(1-ls) - F.logsigmoid(-li)*ls).mean()
print("my dpo b0.5 ls.1", mine.item())
print("ipo", R.ipo_loss(pc,pr,rc,rr,0.5).item())
li=(pc-rc)-(pr-rr)
print("my ipo (li-1/2b)^2", ((li-1/(2*0.5))**2).mean().item())
# grpo_advantages
rewards = rn(6)
print("grpo_adv gs3 std", R.grpo_advantages(rewards,3,True).tolist())
print("grpo_adv gs3 nostd", R.grpo_advantages(rewards,3,False).tolist())
print("rloo_adv gs3", R.rloo_advantages(rewards,3).tolist())
print("group_mean_baseline gs3", R.group_mean_baseline(rewards,3).tolist())
# gae
rw=rn(5); vals=rn(5); nv=rn(1)
print("gae", R.gae(rw,vals,nv.item(),0.99,0.95).tolist() if hasattr(R.gae(rw,vals,nv.item(),0.99,0.95),'tolist') else R.gae(rw,vals,nv.item(),0.99,0.95))
print("lambda_returns", R.lambda_returns(rw,vals,nv.item(),0.99,0.95).tolist())
print("discounted_returns", R.discounted_returns(rw,0.99).tolist())
Run probe 2
python3 p2.py
my masked_mean ent tensor(1.2685) R entropy 1.2684575319290161 R normalized_entropy 0.7881369590759277 norm guess ent/log(V) 0.7881369590759277 dpo b0.1 ls0 0.6615933179855347 dpo b0.5 ls0.1 0.6132094860076904 my dpo b0.1 0.6615933179855347 my dpo b0.5 ls.1 0.6132094860076904 ipo 1.5100568532943726 my ipo (li-1/2b)^2 1.5100568532943726 grpo_adv gs3 std [1.41409432888031, -0.7227914929389954, -0.6913027167320251, 0.3330215811729431, -1.3568118810653687, 1.0237902402877808] grpo_adv gs3 nostd [0.8483808636665344, -0.4336361885070801, -0.41474461555480957, 0.17122069001197815, -0.6975952386856079, 0.5263745188713074] rloo_adv gs3 [1.272571325302124, -0.6504542827606201, -0.6221169233322144, 0.2568310499191284, -1.046392798423767, 0.7895617485046387] group_mean_baseline gs3 [0.8483808636665344, -0.4336361885070801, -0.41474461555480957, 0.17122069001197815, -0.6975952386856079, 0.5263745188713074] gae [-2.528904676437378, -3.6649527549743652, -2.4207582473754883, -1.741837501525879, -2.707975387573242] lambda_returns [-1.2886016368865967, -2.1716012954711914, -1.9220303297042847, -1.5099650621414185, -1.5333479642868042] discounted_returns [-0.47492319345474243, -1.1664562225341797, -0.7856946587562561, -0.27505937218666077, -0.15056981146335602] [stdout] my masked_mean ent tensor(1.2685) R entropy 1.2684575319290161 R normalized_entropy 0.7881369590759277 norm guess ent/log(V) 0.7881369590759277 dpo b0.1 ls0 0.6615933179855347 dpo b0.5 ls0.1 0.6132094860076904 my dpo b0.1 0.6615933179855347 my dpo b0.5 ls.1 0.6132094860076904 ipo 1.5100568532943726 my ipo (li-1/2b)^2 1.5100568532943726 grpo_adv gs3 std [1.41409432888031, -0.7227914929389954, -0.6913027167320251, 0.3330215811729431, -1.3568118810653687, 1.0237902402877808] grpo_adv gs3 nostd [0.8483808636665344, -0.4336361885070801, -0.41474461555480957, 0.17122069001197815, -0.6975952386856079, 0.5263745188713074] rloo_adv gs3 [1.272571325302124, -0.6504542827606201, -0.6221169233322144, 0.2568310499191284, -1.046392798423767, 0.7895617485046387] group_mean_baseline gs3 [0.8483808636665344, -0.4336361885070801, -0.41474461555480957, 0.17122069001197815, -0.6975952386856079, 0.5263745188713074] gae [-2.528904676437378, -3.6649527549743652, -2.4207582473754883, -1.741837501525879, -2.707975387573242] lambda_returns [-1.2886016368865967, -2.1716012954711914, -1.9220303297042847, -1.5099650621414185, -1.5333479642868042] discounted_returns [-0.47492319345474243, -1.1664562225341797, -0.7856946587562561, -0.27505937218666077, -0.15056981146335602]
import torch
import reforacle as R
import torch.nn.functional as F
g = torch.Generator().manual_seed(2)
def rn(*s): return torch.randn(*s, generator=g)
# reproduce p2 rewards for grpo std exact
rewards = None
# reset generator to match p2 sequence up to rewards: emulate
g2=torch.Generator().manual_seed(2)
def rn2(*s): return torch.randn(*s,generator=g2)
_=rn2(2,3,5); _=(torch.rand(2,3,generator=g2)>0.3).float()
pc,pr,rc,rr=rn2(4),rn2(4),rn2(4),rn2(4)
rewards=rn2(6)
print("rewards", rewards.tolist())
grp = rewards.view(2,3)
mean = grp.mean(1,keepdim=True)
for eps in [0,1e-4,1e-6,1e-8]:
for ub in [False,True]:
std = grp.std(1,keepdim=True,unbiased=ub)
adv = ((grp-mean)/(std+eps)).view(-1)
print(f"eps{eps} ub{ub}", adv[0].item())
print("R", R.grpo_advantages(rewards,3,True)[0].item())
# gae exact
gg=torch.Generator().manual_seed(2)
def rn3(*s): return torch.randn(*s,generator=gg)
_=rn3(2,3,5);_=(torch.rand(2,3,generator=gg)>0.3).float()
rn3(4);rn3(4);rn3(4);rn3(4);rn3(6)
rw=rn3(5); vals=rn3(5); nv=rn3(1)
gamma,lam=0.99,0.95
T=5
adv=torch.zeros(T); last=0.0
vnext=nv.item()
for t in reversed(range(T)):
delta=rw[t]+gamma*vnext-vals[t]
last=delta+gamma*lam*last
adv[t]=last
vnext=vals[t]
print("my gae", adv.tolist())
print("R gae", R.gae(rw,vals,nv.item(),gamma,lam).tolist())
print("my lambda", (adv+vals).tolist())
print("R lambda", R.lambda_returns(rw,vals,nv.item(),gamma,lam).tolist())
# discounted
dr=torch.zeros(5); acc=0.0
for t in reversed(range(5)):
acc=rw[t]+0.99*acc; dr[t]=acc
print("my dr", dr.tolist())
print("R dr", R.discounted_returns(rw,0.99).tolist())
Run probe 3
python3 p3.py
rewards [-0.11352437734603882, -1.3955414295196533, -1.3766498565673828, -0.09680592268705368, -0.965621829032898, 0.25834789872169495] eps0 ubFalse 1.414096713066101 eps0 ubTrue 1.1546051502227783 eps0.0001 ubFalse 1.4138610363006592 eps0.0001 ubTrue 1.1544480323791504 eps1e-06 ubFalse 1.41409432888031 eps1e-06 ubTrue 1.1546036005020142 eps1e-08 ubFalse 1.414096713066101 eps1e-08 ubTrue 1.1546051502227783 R 1.41409432888031 my gae [-2.528904676437378, -3.6649527549743652, -2.4207582473754883, -1.741837501525879, -2.707975387573242] R gae [-2.528904676437378, -3.6649527549743652, -2.4207582473754883, -1.741837501525879, -2.707975387573242] my lambda [-1.2886016368865967, -2.1716012954711914, -1.9220303297042847, -1.5099650621414185, -1.5333479642868042] R lambda [-1.2886016368865967, -2.1716012954711914, -1.9220303297042847, -1.5099650621414185, -1.5333479642868042] my dr [-0.47492319345474243, -1.1664562225341797, -0.7856946587562561, -0.27505937218666077, -0.15056981146335602] R dr [-0.47492319345474243, -1.1664562225341797, -0.7856946587562561, -0.27505937218666077, -0.15056981146335602] [stdout] rewards [-0.11352437734603882, -1.3955414295196533, -1.3766498565673828, -0.09680592268705368, -0.965621829032898, 0.25834789872169495] eps0 ubFalse 1.414096713066101 eps0 ubTrue 1.1546051502227783 eps0.0001 ubFalse 1.4138610363006592 eps0.0001 ubTrue 1.1544480323791504 eps1e-06 ubFalse 1.41409432888031 eps1e-06 ubTrue 1.1546036005020142 eps1e-08 ubFalse 1.414096713066101 eps1e-08 ubTrue 1.1546051502227783 R 1.41409432888031 my gae [-2.528904676437378, -3.6649527549743652, -2.4207582473754883, -1.741837501525879, -2.707975387573242] R gae [-2.528904676437378, -3.6649527549743652, -2.4207582473754883, -1.741837501525879, -2.707975387573242] my lambda [-1.2886016368865967, -2.1716012954711914, -1.9220303297042847, -1.5099650621414185, -1.5333479642868042] R lambda [-1.2886016368865967, -2.1716012954711914, -1.9220303297042847, -1.5099650621414185, -1.5333479642868042] my dr [-0.47492319345474243, -1.1664562225341797, -0.7856946587562561, -0.27505937218666077, -0.15056981146335602] R dr [-0.47492319345474243, -1.1664562225341797, -0.7856946587562561, -0.27505937218666077, -0.15056981146335602]
import torch
import reforacle as R
import torch.nn.functional as F
g = torch.Generator().manual_seed(5)
def rn(*s): return torch.randn(*s, generator=g)
logp=rn(3,4); ref=rn(3,4); old=rn(3,4)
print("reverse_kl", R.reverse_kl(logp,ref).tolist() if hasattr(R.reverse_kl(logp,ref),'tolist') else R.reverse_kl(logp,ref))
print(" my logp-ref mean?", (logp-ref).mean().item(), "elementwise?")
rk=R.reverse_kl(logp,ref)
print(" rk shape", rk.shape if hasattr(rk,'shape') else 'scalar')
print("symmetric_kl", R.symmetric_kl(logp,ref))
sk=R.symmetric_kl(logp,ref)
print(" sk shape", sk.shape if hasattr(sk,'shape') else 'scalar')
# kl_penalty estimators
for est in ['k1','k2','k3','mse','kl','abs']:
try:
r=R.kl_penalty(logp,ref,est)
print("kl_penalty",est, r.shape if hasattr(r,'shape') else r, (r.flatten()[:3].tolist() if hasattr(r,'flatten') else r))
except Exception as e:
print("kl_penalty",est,"ERR",e)
# my guesses
d=logp-ref
print("k1=logp-ref", d.flatten()[:3].tolist())
print("k2=.5*d^2", (0.5*d*d).flatten()[:3].tolist())
print("k3=exp(-d)-1+d? actually ref-logp form", (torch.exp(ref-logp)-1-(ref-logp)).flatten()[:3].tolist())
# importance_ratio
print("importance_ratio noclip", R.importance_ratio(logp,old,None))
ir=R.importance_ratio(logp,old,None)
print(" exp(logp-old)", torch.exp(logp-old).flatten()[:3].tolist(), ir.flatten()[:3].tolist() if hasattr(ir,'flatten') else ir)
print("importance_ratio clip0.2", R.importance_ratio(logp,old,0.2).flatten()[:3].tolist())
# clip_fraction
print("clip_fraction", R.clip_fraction(logp,old,0.2))
# normalize
x=rn(6)
print("normalize eps1e-8", R.normalize(x,1e-8).tolist())
print(" my", ((x-x.mean())/(x.std()+1e-8)).tolist(), "unbiased?")
print(" my ub F", ((x-x.mean())/(x.std(unbiased=False)+1e-8)).tolist())
# bradley_terry_logit
print("bt", R.bradley_terry_logit(torch.tensor(1.0),torch.tensor(0.3),0.5))
Run probe 4
python3 p4.py
reverse_kl [[0.3232053518295288, 0.0019072294235229492, 1.22235107421875, 0.0042400360107421875], [0.06560015678405762, 0.4837096929550171, 1.8419780731201172, 0.07940006256103516], [0.49916768074035645, 0.0022597312927246094, 1.8434550762176514, 0.09845864772796631]]
my logp-ref mean? -0.09309972077608109 elementwise?
rk shape torch.Size([3, 4])
symmetric_kl tensor([[0.2622, 0.0019, 0.8768, 0.0044],
[0.0752, 0.7732, 1.2666, 0.0925],
[0.8063, 0.0022, 1.2676, 0.1170]])
sk shape torch.Size([3, 4])
kl_penalty k1 torch.Size([3, 4]) [-0.7092393636703491, 0.06240367889404297, -1.2427384853363037]
kl_penalty k2 torch.Size([3, 4]) [0.25151023268699646, 0.0019471095874905586, 0.7721994519233704]
kl_penalty k3 torch.Size([3, 4]) [0.3232053518295288, 0.0019072294235229492, 1.22235107421875]
kl_penalty mse ERR mse
kl_penalty kl ERR kl
kl_penalty abs ERR abs
k1=logp-ref [-0.7092393636703491, 0.06240367889404297, -1.2427384853363037]
k2=.5*d^2 [0.25151023268699646, 0.0019471095874905586, 0.7721994519233704]
k3=exp(-d)-1+d? actually ref-logp form [0.3232053518295288, 0.0019072294235229492, 1.22235107421875]
importance_ratio noclip tensor([[ 0.3927, 1.0050, 0.4507, 1.2185],
[ 5.7926, 12.2010, 0.1602, 4.4896],
[ 2.3000, 0.1399, 0.4138, 2.1363]])
exp(logp-old) [0.39270028471946716, 1.0050145387649536, 0.450659841299057] [0.39270028471946716, 1.0050145387649536, 0.450659841299057]
importance_ratio clip0.2 [0.800000011920929, 1.0050145387649536, 0.800000011920929]
clip_fraction tensor(0.9167)
normalize eps1e-8 [-0.8491652607917786, 1.2372227907180786, -0.2706614136695862, 1.294490098953247, -0.6704100370407104, -0.7414760589599609]
my [-0.8491652607917786, 1.2372227907180786, -0.2706614136695862, 1.294490098953247, -0.6704100370407104, -0.7414760589599609] unbiased?
my ub F [-0.9302139282226562, 1.3553096055984497, -0.2964946925640106, 1.418042778968811, -0.7343973517417908, -0.8122462630271912]
bt tensor(0.3500)
[stdout]
reverse_kl [[0.3232053518295288, 0.0019072294235229492, 1.22235107421875, 0.0042400360107421875], [0.06560015678405762, 0.4837096929550171, 1.8419780731201172, 0.07940006256103516], [0.49916768074035645, 0.0022597312927246094, 1.8434550762176514, 0.09845864772796631]]
my logp-ref mean? -0.09309972077608109 elementwise?
rk shape torch.Size([3, 4])
symmetric_kl tensor([[0.2622, 0.0019, 0.8768, 0.0044],
[0.0752, 0.7732, 1.2666, 0.0925],
[0.8063, 0.0022, 1.2676, 0.1170]])
sk shape torch.Size([3, 4])
kl_penalty k1 torch.Size([3, 4]) [-0.7092393636703491, 0.06240367889404297, -1.2427384853363037]
kl_penalty k2 torch.Size([3, 4]) [0.25151023268699646, 0.0019471095874905586, 0.7721994519233704]
kl_penalty k3 torch.Size([3, 4]) [0.3232053518295288, 0.0019072294235229492, 1.22235107421875]
kl_penalty mse ERR mse
kl_penalty kl ERR kl
kl_penalty abs ERR abs
k1=logp-ref [-0.7092393636703491, 0.06240367889404297, -1.2427384853363037]
k2=.5*d^2 [0.25151023268699646, 0.0019471095874905586, 0.7721994519233704]
k3=exp(-d)-1+d? actually ref-logp form [0.3232053518295288, 0.0019072294235229492, 1.22235107421875]
importance_ratio noclip tensor([[ 0.3927, 1.0050, 0.4507, 1.2185],
[ 5.7926, 12.2010, 0.1602, 4.4896],
[ 2.3000, 0.1399, 0.4138, 2.1363]])
exp(logp-old) [0.39270028471946716, 1.0050145387649536, 0.450659841299057] [0.39270028471946716, 1.0050145387649536, 0.450659841299057]
importance_ratio clip0.2 [0.800000011920929, 1.0050145387649536, 0.800000011920929]
clip_fraction tensor(0.9167)
normalize eps1e-8 [-0.8491652607917786, 1.2372227907180786, -0.2706614136695862, 1.294490098953247, -0.6704100370407104, -0.7414760589599609]
my [-0.8491652607917786, 1.2372227907180786, -0.2706614136695862, 1.294490098953247, -0.6704100370407104, -0.7414760589599609] unbiased?
my ub F [-0.9302139282226562, 1.3553096055984497, -0.2964946925640106, 1.418042778968811, -0.7343973517417908, -0.8122462630271912]
bt tensor(0.3500)import torch
import reforacle as R
import torch.nn.functional as F
g = torch.Generator().manual_seed(7)
def rn(*s): return torch.randn(*s, generator=g)
logp=rn(2,4); old=rn(2,4); adv=rn(2,4); mask=(torch.rand(2,4,generator=g)>0.3).float()
r=R.clipped_pg_loss(logp,old,adv,mask,0.2,0.2)
print("clipped_pg_loss", r)
# my: ratio=exp(logp-old); l1=ratio*adv; l2=clamp(ratio,1-lo,1+hi)*adv; loss=-min(l1,l2); masked_mean
ratio=torch.exp(logp-old)
l1=ratio*adv; l2=torch.clamp(ratio,1-0.2,1+0.2)*adv
loss=-torch.min(l1,l2)
print("my masked_mean", ((loss*mask).sum()/mask.sum()).item())
# value_loss
values=rn(2,4); oldv=rn(2,4); returns=rn(2,4)
print("value_loss clip0.2", R.value_loss(values,oldv,returns,0.2))
vc=oldv+torch.clamp(values-oldv,-0.2,0.2)
vl1=(values-returns)**2; vl2=(vc-returns)**2
print("my max mean", (0.5*torch.max(vl1,vl2)).mean().item())
print("my max mean nohalf", (torch.max(vl1,vl2)).mean().item())
# whiten
x=rn(10)
print("whiten shift True", R.whiten(x,None,True) if False else None)
print("whiten shiftT", R.whiten(x, True).tolist() if False else "skip")
Run probe 5
python3 p5.py
clipped_pg_loss tensor(-0.6341) my masked_mean -0.634052574634552 value_loss clip0.2 tensor(0.6550) my max mean 0.6549816727638245 my max mean nohalf 1.309963345527649 whiten shift True None whiten shiftT skip [stdout] clipped_pg_loss tensor(-0.6341) my masked_mean -0.634052574634552 value_loss clip0.2 tensor(0.6550) my max mean 0.6549816727638245 my max mean nohalf 1.309963345527649 whiten shift True None whiten shiftT skip
import torch
import reforacle as R
g = torch.Generator().manual_seed(9)
def rn(*s): return torch.randn(*s, generator=g)
x=rn(10)
mask=(torch.rand(10,generator=g)>0.3).float()
print("whiten mask shiftT", R.whiten(x,mask,True).tolist())
print("whiten mask shiftF", R.whiten(x,mask,False).tolist())
print("whiten None shiftT", R.whiten(x,None,True).tolist())
# guess whiten: uses mask? or ignores? shift_mean subtract mean
for ub in [True,False]:
m=x.mean(); s=x.std(unbiased=ub)
print(f"my full ub{ub} shiftT", ((x-m)/(s+1e-8)).tolist()[:3])
print(f"my full ub{ub} shiftF", ((x)/(s+1e-8)).tolist()[:3])
print("masked_whiten shiftT", R.masked_whiten(x,mask,True).tolist())
print("masked_whiten shiftF", R.masked_whiten(x,mask,False).tolist())
Run probe 6
python3 p6.py
Exit code 1
whiten mask shiftT [-0.9576643109321594, 1.966176152229309, -1.3896325826644897, -0.4858416020870209, 1.0369272232055664, -0.448702335357666, -0.4260283410549164, -0.5624263286590576, 0.14287561178207397, 0.675614595413208]
whiten mask shiftF [-0.30158621072769165, 2.6222543716430664, -0.733554482460022, 0.1702364981174469, 1.6930053234100342, 0.20737576484680176, 0.2300497591495514, 0.09365177154541016, 0.7989537119865417, 1.3316926956176758]
Traceback (most recent call last):
File "/app/p6.py", line 10, in <module>
print("whiten None shiftT", R.whiten(x,None,True).tolist())
^^^^^^^^^^^^^^^^^^^^^
File "reforacle.py", line 115, in reforacle.whiten
File "reforacle.py", line 10, in reforacle.masked_mean
AttributeError: 'NoneType' object has no attribute 'to'
[error] tool reported failureimport torch
import reforacle as R
g = torch.Generator().manual_seed(9)
def rn(*s): return torch.randn(*s, generator=g)
x=rn(10)
mask=(torch.rand(10,generator=g)>0.3).float()
def mmean(v,m): return (v*m).sum()/m.sum()
mean=mmean(x,mask)
# masked var: population vs unbiased
n=mask.sum()
var_pop=mmean((x-mean)**2,mask)
var_ub = ((x-mean)**2*mask).sum()/(n-1)
for name,var in [("pop",var_pop),("ub",var_ub)]:
for eps in [1e-8,1e-6,0]:
w=(x-mean)/torch.sqrt(var+eps)
print(f"shiftT {name} eps{eps}", w[:3].tolist())
print("R whiten shiftT", R.whiten(x,mask,True)[:3].tolist())
print("R masked_whiten shiftT", R.masked_whiten(x,mask,True)[:3].tolist())
# shiftF: add mean back? whiten typically: whitened = (x-mean)/std; if not shift_mean: += mean
wpop=(x-mean)/torch.sqrt(var_pop+1e-8)
print("shiftF = w+mean", (wpop+mean)[:3].tolist())
print("R whiten shiftF", R.whiten(x,mask,False)[:3].tolist())
print("R masked_whiten shiftF", R.masked_whiten(x,mask,False)[:3].tolist())
Run probe 7
python3 p7.py
shiftT pop eps1e-08 [-0.9576643109321594, 1.966176152229309, -1.3896325826644897] shiftT pop eps1e-06 [-0.9576631188392639, 1.966173768043518, -1.389630913734436] shiftT pop eps0 [-0.9576643109321594, 1.966176152229309, -1.3896325826644897] shiftT ub eps1e-08 [-0.9028945565223694, 1.8537285327911377, -1.310158133506775] shiftT ub eps1e-06 [-0.902893602848053, 1.8537266254425049, -1.3101567029953003] shiftT ub eps0 [-0.9028945565223694, 1.8537285327911377, -1.310158133506775] R whiten shiftT [-0.9576643109321594, 1.966176152229309, -1.3896325826644897] R masked_whiten shiftT [-0.9576643109321594, 1.966176152229309, -1.3896325826644897] shiftF = w+mean [-0.30158621072769165, 2.6222543716430664, -0.733554482460022] R whiten shiftF [-0.30158621072769165, 2.6222543716430664, -0.733554482460022] R masked_whiten shiftF [-0.30158621072769165, 2.6222543716430664, -0.733554482460022] [stdout] shiftT pop eps1e-08 [-0.9576643109321594, 1.966176152229309, -1.3896325826644897] shiftT pop eps1e-06 [-0.9576631188392639, 1.966173768043518, -1.389630913734436] shiftT pop eps0 [-0.9576643109321594, 1.966176152229309, -1.3896325826644897] shiftT ub eps1e-08 [-0.9028945565223694, 1.8537285327911377, -1.310158133506775] shiftT ub eps1e-06 [-0.902893602848053, 1.8537266254425049, -1.3101567029953003] shiftT ub eps0 [-0.9028945565223694, 1.8537285327911377, -1.310158133506775] R whiten shiftT [-0.9576643109321594, 1.966176152229309, -1.3896325826644897] R masked_whiten shiftT [-0.9576643109321594, 1.966176152229309, -1.3896325826644897] shiftF = w+mean [-0.30158621072769165, 2.6222543716430664, -0.733554482460022] R whiten shiftF [-0.30158621072769165, 2.6222543716430664, -0.733554482460022] R masked_whiten shiftF [-0.30158621072769165, 2.6222543716430664, -0.733554482460022]
import torch
import reforacle as R
import torch.nn.functional as F
g = torch.Generator().manual_seed(11)
def rn(*s): return torch.randn(*s, generator=g)
# cross_entropy
logits=rn(2,3,5); labels=torch.randint(0,5,(2,3),generator=g)
labels[0,1]=-100
print("cross_entropy ii-100", R.cross_entropy(logits,labels,-100))
print(" F.cross_entropy mean", F.cross_entropy(logits.reshape(-1,5),labels.reshape(-1),ignore_index=-100).item())
ce=R.cross_entropy(logits,labels,-100)
print(" shape", ce.shape if hasattr(ce,'shape') else 'scalar')
# smoothed_nll
labels2=torch.randint(0,5,(2,3),generator=g)
print("smoothed_nll 0.1", R.smoothed_nll(logits,labels2,0.1))
sm=R.smoothed_nll(logits,labels2,0.1)
print(" shape", sm.shape if hasattr(sm,'shape') else 'scalar')
lp=torch.log_softmax(logits,-1)
nll=-lp.gather(-1,labels2.unsqueeze(-1)).squeeze(-1)
smooth=-lp.mean(-1)
s=0.1
print(" my (1-s)nll+s*smooth mean", ((1-s)*nll+s*smooth).mean().item())
print(" my sum?", ((1-s)*nll+s*smooth).sum().item())
# logprob_at_temperature
print("logprob_at_temp T2", R.logprob_at_temperature(logits,labels2,2.0))
lt=R.logprob_at_temperature(logits,labels2,2.0)
print(" shape", lt.shape if hasattr(lt,'shape') else 'scalar')
lpt=torch.log_softmax(logits/2.0,-1)
print(" my gather", lpt.gather(-1,labels2.unsqueeze(-1)).squeeze(-1).tolist())
# huber_value_loss
values=rn(6); returns=rn(6)
print("huber d1", R.huber_value_loss(values,returns,1.0))
print(" F.huber mean", F.huber_loss(values,returns,delta=1.0).item())
print(" F.huber sum", F.huber_loss(values,returns,delta=1.0,reduction='sum').item())
# advantage_mean_std
adv=rn(8); mask=(torch.rand(8,generator=g)>0.3).float()
print("advantage_mean_std", R.advantage_mean_std(adv,mask))
Run probe 8
python3 p8.py
cross_entropy ii-100 tensor(1.5063)
F.cross_entropy mean 1.506290078163147
shape torch.Size([])
smoothed_nll 0.1 tensor(2.2460)
shape torch.Size([])
my (1-s)nll+s*smooth mean 2.246000289916992
my sum? 13.476001739501953
logprob_at_temp T2 tensor([[-1.1463, -1.6404, -2.4385],
[-1.7526, -1.7045, -2.2561]])
shape torch.Size([2, 3])
my gather [[-1.1462960243225098, -1.6404414176940918, -2.438537120819092], [-1.7525713443756104, -1.704535722732544, -2.2561211585998535]]
huber d1 tensor(0.2609)
F.huber mean 0.260940819978714
F.huber sum 1.5656449794769287
advantage_mean_std tensor([0.0078, 0.5737])
[stdout]
cross_entropy ii-100 tensor(1.5063)
F.cross_entropy mean 1.506290078163147
shape torch.Size([])
smoothed_nll 0.1 tensor(2.2460)
shape torch.Size([])
my (1-s)nll+s*smooth mean 2.246000289916992
my sum? 13.476001739501953
logprob_at_temp T2 tensor([[-1.1463, -1.6404, -2.4385],
[-1.7526, -1.7045, -2.2561]])
shape torch.Size([2, 3])
my gather [[-1.1462960243225098, -1.6404414176940918, -2.438537120819092], [-1.7525713443756104, -1.704535722732544, -2.2561211585998535]]
huber d1 tensor(0.2609)
F.huber mean 0.260940819978714
F.huber sum 1.5656449794769287
advantage_mean_std tensor([0.0078, 0.5737])import torch
import reforacle as R
g = torch.Generator().manual_seed(11)
def rn(*s): return torch.randn(*s, generator=g)
# replicate p8 adv/mask
gg=torch.Generator().manual_seed(11)
def r2(*s): return torch.randn(*s,generator=gg)
r2(2,3,5); torch.randint(0,5,(2,3),generator=gg); torch.randint(0,5,(2,3),generator=gg)
r2(6); r2(6); adv=r2(8); mask=(torch.rand(8,generator=gg)>0.3).float()
def mmean(v,m):return (v*m).sum()/m.sum()
mean=mmean(adv,mask); n=mask.sum()
var_pop=mmean((adv-mean)**2,mask)
var_ub=((adv-mean)**2*mask).sum()/(n-1)
print("R", R.advantage_mean_std(adv,mask).tolist())
print("mean",mean.item(),"std_pop",var_pop.sqrt().item(),"std_ub",var_ub.sqrt().item())
g = torch.Generator().manual_seed(13)
def rn(*s): return torch.randn(*s, generator=g)
# top_k_mask
logits=rn(2,6)
print("top_k_mask k3", R.top_k_mask(logits,3))
# top_p_mask
probs=torch.softmax(rn(2,6),-1)
print("probs", probs.tolist())
print("top_p_mask 0.8", R.top_p_mask(probs,0.8))
# argmax_tokens
l2=rn(2,3,5)
print("argmax_tokens", R.argmax_tokens(l2), R.argmax_tokens(l2).dtype)
# mode_label
labs=torch.tensor([1,2,2,3,2,1])
print("mode_label", R.mode_label(labs))
# median_reward
rw=rn(7)
print("rw",rw.tolist())
print("median_reward", R.median_reward(rw))
print("torch.median", torch.median(rw).item(), "quantile.5", torch.quantile(rw,0.5).item())
# masked_max
v=rn(8); m=(torch.rand(8,generator=g)>0.3).float()
print("v",v.tolist(),"m",m.tolist())
print("masked_max", R.masked_max(v,m))
# quantile_lower
x=rn(10)
print("quantile_lower 0.25", R.quantile_lower(x,0.25))
print("torch.quantile lower", torch.quantile(x,0.25,interpolation='lower').item(), "linear", torch.quantile(x,0.25).item())
# pad_mask_from_lengths
print("pad_mask", R.pad_mask_from_lengths(torch.tensor([2,4,1]),5))
# first_nonzero_index
mm=torch.tensor([[0.,0,1,0,1],[0,0,0,0,0],[1,0,0,0,0]])
print("first_nonzero", R.first_nonzero_index(mm))
# cumulative_max
print("cumulative_max", R.cumulative_max(torch.tensor([1.,3,2,5,4])))
# bucketize_reward
print("bucketize", R.bucketize_reward(torch.tensor([-1.,0.5,1.5,3.0]), torch.tensor([0.,1.,2.])))
Run probe 9
python3 p9.py
R [0.007819448597729206, 0.573697566986084]
mean 0.007819448597729206 std_pop 0.573697566986084 std_ub 0.619664192199707
top_k_mask k3 tensor([[False, False, True, True, True, False],
[False, True, True, True, False, False]])
probs [[0.21831350028514862, 0.07545264065265656, 0.07304232567548752, 0.08553174883127213, 0.5140340924263, 0.03362565487623215], [0.048587214201688766, 0.49695298075675964, 0.06684879958629608, 0.03170686587691307, 0.19017952680587769, 0.1657245010137558]]
top_p_mask 0.8 tensor([[0.2669, 0.0000, 0.0000, 0.1046, 0.6285, 0.0000],
[0.0000, 0.5827, 0.0000, 0.0000, 0.2230, 0.1943]])
argmax_tokens tensor([[2, 3, 3],
[0, 4, 3]]) torch.int64
mode_label tensor(2)
rw [1.5578457117080688, 1.0663607120513916, 0.3840605318546295, 0.261141300201416, 0.6970183849334717, -1.1506617069244385, 0.04642544686794281]
median_reward tensor(0.3841)
torch.median 0.3840605318546295 quantile.5 0.3840605318546295
v [2.1695189476013184, 0.6781666278839111, -0.8294501900672913, 0.9859806895256042, -1.270143747329712, 2.51708984375, 1.3225687742233276, 1.5655877590179443] m [1.0, 1.0, 0.0, 1.0, 1.0, 1.0, 1.0, 1.0]
masked_max tensor(2.5171)
quantile_lower 0.25 tensor(-0.3061)
torch.quantile lower -0.30611497163772583 linear -0.29328420758247375
pad_mask tensor([[1., 1., 0., 0., 0.],
[1., 1., 1., 1., 0.],
[1., 0., 0., 0., 0.]])
first_nonzero tensor([ 2, -1, 0])
cumulative_max tensor([1., 1., 3., 3., 5.])
bucketize tensor([0, 1, 2, 3])
[stdout]
R [0.007819448597729206, 0.573697566986084]
mean 0.007819448597729206 std_pop 0.573697566986084 std_ub 0.619664192199707
top_k_mask k3 tensor([[False, False, True, True, True, False],
[False, True, True, True, False, False]])
probs [[0.21831350028514862, 0.07545264065265656, 0.07304232567548752, 0.08553174883127213, 0.5140340924263, 0.03362565487623215], [0.048587214201688766, 0.49695298075675964, 0.06684879958629608, 0.03170686587691307, 0.19017952680587769, 0.1657245010137558]]
top_p_mask 0.8 tensor([[0.2669, 0.0000, 0.0000, 0.1046, 0.6285, 0.0000],
[0.0000, 0.5827, 0.0000, 0.0000, 0.2230, 0.1943]])
argmax_tokens tensor([[2, 3, 3],
[0, 4, 3]]) torch.int64
mode_label tensor(2)
rw [1.5578457117080688, 1.0663607120513916, 0.3840605318546295, 0.261141300201416, 0.6970183849334717, -1.1506617069244385, 0.04642544686794281]
median_reward tensor(0.3841)
torch.median 0.3840605318546295 quantile.5 0.3840605318546295
v [2.1695189476013184, 0.6781666278839111, -0.8294501900672913, 0.9859806895256042, -1.270143747329712, 2.51708984375, 1.3225687742233276, 1.5655877590179443] m [1.0, 1.0, 0.0, 1.0, 1.0, 1.0, 1.0, 1.0]
masked_max tensor(2.5171)
quantile_lower 0.25 tensor(-0.3061)
torch.quantile lower -0.30611497163772583 linear -0.29328420758247375
pad_mask tensor([[1., 1., 0., 0., 0.],
[1., 1., 1., 1., 0.],
[1., 0., 0., 0., 0.]])
first_nonzero tensor([ 2, -1, 0])
cumulative_max tensor([1., 1., 3., 3., 5.])
bucketize tensor([0, 1, 2, 3])import torch
import reforacle as R
import torch.nn.functional as F
g = torch.Generator().manual_seed(21)
def rn(*s): return torch.randn(*s, generator=g)
V=5;B=2;T=4
pcl=rn(B,T,V);prl=rn(B,T,V);rcl=rn(B,T,V);rrl=rn(B,T,V)
cl=torch.randint(0,V,(B,T),generator=g); rl=torch.randint(0,V,(B,T),generator=g)
cm=(torch.rand(B,T,generator=g)>0.3).float(); rm=(torch.rand(B,T,generator=g)>0.3).float()
beta=0.1; ls=0.0
r=R.dpo_sequence_loss(pcl,prl,rcl,rrl,cl,rl,cm,rm,beta,ls)
print("dpo_seq", r)
# my
def seqlp(logits,labels,mask):
lp=torch.log_softmax(logits,-1)
tl=lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
return (tl*mask).sum(-1)
pc=seqlp(pcl,cl,cm);pr=seqlp(prl,rl,rm);rc=seqlp(rcl,cl,cm);rr=seqlp(rrl,rl,rm)
li=beta*((pc-rc)-(pr-rr))
mine=(-F.logsigmoid(li)*(1-ls)-F.logsigmoid(-li)*ls).mean()
print("my dpo_seq", mine.item())
print("---rloo_objective---")
# rloo_objective(logits, old_logits, labels, mask, rewards, group_size, clip_low, clip_high)
G=2; N=4
logits=rn(N,T,V); old=rn(N,T,V); labels=torch.randint(0,V,(N,T),generator=g)
mask=(torch.rand(N,T,generator=g)>0.3).float(); rewards=rn(N)
r=R.rloo_objective(logits,old,labels,mask,rewards,G,0.2,0.2)
print("rloo_obj", r)
Run probe 10
python3 p10.py
dpo_seq tensor(0.8746) my dpo_seq 0.8746036291122437 ---rloo_objective--- rloo_obj tensor(0.0667) [stdout] dpo_seq tensor(0.8746) my dpo_seq 0.8746036291122437 ---rloo_objective--- rloo_obj tensor(0.0667)
import torch
import reforacle as R
import torch.nn.functional as F
g = torch.Generator().manual_seed(21)
def rn(*s): return torch.randn(*s, generator=g)
V=5;B=2;T=4
# replicate up to rloo inputs
rn(B,T,V);rn(B,T,V);rn(B,T,V);rn(B,T,V)
torch.randint(0,V,(B,T),generator=g);torch.randint(0,V,(B,T),generator=g)
(torch.rand(B,T,generator=g)>0.3).float();(torch.rand(B,T,generator=g)>0.3).float()
G=2;N=4
logits=rn(N,T,V);old=rn(N,T,V);labels=torch.randint(0,V,(N,T),generator=g)
mask=(torch.rand(N,T,generator=g)>0.3).float();rewards=rn(N)
def tlp(lg,lb):
lp=torch.log_softmax(lg,-1); return lp.gather(-1,lb.unsqueeze(-1)).squeeze(-1)
logp=tlp(logits,labels); oldlp=tlp(old,labels)
adv=R.rloo_advantages(rewards,G) # (N,)
ratio=torch.exp(logp-oldlp)
a=adv.unsqueeze(-1)
l1=ratio*a; l2=torch.clamp(ratio,0.8,1.2)*a
loss=-torch.min(l1,l2)
print("masked_mean all", ((loss*mask).sum()/mask.sum()).item())
# per-seq masked mean then mean
per=(loss*mask).sum(-1)/mask.sum(-1)
print("per-seq then mean", per.mean().item())
print("R", R.rloo_objective(logits,old,labels,mask,rewards,G,0.2,0.2).item())
Run probe 11
python3 p11.py
masked_mean all 0.06669013202190399 per-seq then mean 0.0511140376329422 R 0.06669013202190399 [stdout] masked_mean all 0.06669013202190399 per-seq then mean 0.0511140376329422 R 0.06669013202190399
import torch
import reforacle as R
import torch.nn.functional as F
g = torch.Generator().manual_seed(31)
def rn(*s): return torch.randn(*s, generator=g)
V=5;N=3;T=4
logits=rn(N,T,V); ref=rn(N,T,V); labels=torch.randint(0,V,(N,T),generator=g)
mask=(torch.rand(N,T,generator=g)>0.3).float(); adv=rn(N,T); beta=0.1
R_val=R.reverse_kl_objective(logits,ref,labels,mask,adv,beta).item()
print("R", R_val)
lp=torch.log_softmax(logits,-1); rlp=torch.log_softmax(ref,-1)
tl=lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
rtl=rlp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
# pg part
pg=-(adv*tl)
# reverse_kl token estimator k3 on token logps: exp(rtl-tl)-1-(rtl-tl)
d=tl-rtl
kl_tok=torch.exp(-d)-1+d
def mm(x): return (x*mask).sum()/mask.sum()
print("cand1 pg+beta*kl (k3 token)", (mm(pg)+beta*mm(kl_tok)).item())
# full dist reverse kl: sum_v p*(logp-logref)
full_kl=(lp.exp()*(lp-rlp)).sum(-1)
print("cand2 pg+beta*full_kl", (mm(pg)+beta*mm(full_kl)).item())
# maybe advantage-weighted with adv per seq? here adv is (N,T)
# maybe loss = mm(-adv*tl + beta*kl)
print("cand3 mm(pg + beta*kl_tok)", mm(pg+beta*kl_tok).item())
print("cand4 mm(pg + beta*full_kl)", mm(pg+beta*full_kl).item())
Run probe 12
python3 p12.py
R 1.4292851686477661 cand1 pg+beta*kl (k3 token) 0.41307511925697327 cand2 pg+beta*full_kl 0.43341895937919617 cand3 mm(pg + beta*kl_tok) 0.41307514905929565 cand4 mm(pg + beta*full_kl) 0.43341895937919617 [stdout] R 1.4292851686477661 cand1 pg+beta*kl (k3 token) 0.41307511925697327 cand2 pg+beta*full_kl 0.43341895937919617 cand3 mm(pg + beta*kl_tok) 0.41307514905929565 cand4 mm(pg + beta*full_kl) 0.43341895937919617
import torch
import reforacle as R
import torch.nn.functional as F
g = torch.Generator().manual_seed(31)
def rn(*s): return torch.randn(*s, generator=g)
V=5;N=3;T=4
logits=rn(N,T,V); ref=rn(N,T,V); labels=torch.randint(0,V,(N,T),generator=g)
mask=(torch.rand(N,T,generator=g)>0.3).float(); adv=rn(N,T); beta=0.1
lp=torch.log_softmax(logits,-1); rlp=torch.log_softmax(ref,-1)
tl=lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
rtl=rlp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
d=tl-rtl; kl_tok=torch.exp(-d)-1+d
def mm(x): return ((x*mask).sum()/mask.sum()).item()
print("mm(-adv*tl)",mm(-adv*tl))
print("mm(adv*tl)",mm(adv*tl))
print("mm(kl_tok)",mm(kl_tok))
print("mm(-tl)",mm(-tl))
print("beta*mm(kl)",0.1*mm(kl_tok))
# R=1.4293. Try: -mm(adv*tl)+beta*mm(kl) with different:
print("-mm(adv*tl)+beta*mm(kl)", -mm(adv*tl)+0.1*mm(kl_tok))
# ratio-based pg with old=ref?
ratio=torch.exp(tl-rtl)
print("mm(-ratio*adv)+beta*kl", mm(-ratio*adv)+0.1*mm(kl_tok))
print("mm(-ratio*adv)", mm(-ratio*adv))
# clamp? no clip args.
# maybe advantage applied and kl not masked_mean but masked_sum
def ms(x): return ((x*mask).sum()).item()
print("R again", R.reverse_kl_objective(logits,ref,labels,mask,adv,beta).item())
Run probe 13
python3 p13.py
mm(-adv*tl) 0.3551923930644989 mm(adv*tl) -0.3551923930644989 mm(kl_tok) 0.5788273215293884 mm(-tl) 1.958479642868042 beta*mm(kl) 0.05788273215293885 -mm(adv*tl)+beta*mm(kl) 0.4130751252174377 mm(-ratio*adv)+beta*kl -0.01764689534902572 mm(-ratio*adv) -0.07552962750196457 R again 1.4292851686477661 [stdout] mm(-adv*tl) 0.3551923930644989 mm(adv*tl) -0.3551923930644989 mm(kl_tok) 0.5788273215293884 mm(-tl) 1.958479642868042 beta*mm(kl) 0.05788273215293885 -mm(adv*tl)+beta*mm(kl) 0.4130751252174377 mm(-ratio*adv)+beta*kl -0.01764689534902572 mm(-ratio*adv) -0.07552962750196457 R again 1.4292851686477661
import torch, itertools
import reforacle as R
g = torch.Generator().manual_seed(31)
def rn(*s): return torch.randn(*s, generator=g)
V=5;N=3;T=4
logits=rn(N,T,V); ref=rn(N,T,V); labels=torch.randint(0,V,(N,T),generator=g)
mask=(torch.rand(N,T,generator=g)>0.3).float(); adv=rn(N,T); beta=0.1
lp=torch.log_softmax(logits,-1); rlp=torch.log_softmax(ref,-1)
tl=lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
rtl=rlp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
d=tl-rtl
def mm(x): return ((x*mask).sum()/mask.sum())
kl_tok=torch.exp(-d)-1+d
full_rev=(lp.exp()*(lp-rlp)).sum(-1)
full_fwd=(rlp.exp()*(rlp-lp)).sum(-1)
ratio=torch.exp(d)
blocks={
'adv*tl':mm(adv*tl),'adv*d':mm(adv*d),'tl':mm(tl),'d':mm(d),
'kl_tok':mm(kl_tok),'full_rev':mm(full_rev),'full_fwd':mm(full_fwd),
'ratio*adv':mm(ratio*adv),'adv':mm(adv)
}
for k,v in blocks.items(): print(k, v.item())
print("TARGET", R.reverse_kl_objective(logits,ref,labels,mask,adv,beta).item())
# search a*block1 + b*block2 with a in signs, plus beta*kl
import numpy as np
target=1.4292851686477661
names=list(blocks); vals=[blocks[n].item() for n in names]
best=[]
for i in range(len(names)):
for si in [1,-1]:
base=si*vals[i]
# + beta*kl_tok or +beta*full_rev or nothing
for extra,ev in [('none',0),('b*kl_tok',0.1*blocks['kl_tok'].item()),('b*full_rev',0.1*blocks['full_rev'].item())]:
if abs(base+ev-target)<1e-4:
best.append((si,names[i],extra))
print("MATCH1", best)
# two-block combos with unit coeffs
best2=[]
for i in range(len(names)):
for j in range(len(names)):
for si in[1,-1]:
for sj in[1,-1]:
val=si*vals[i]+sj*vals[j]
if abs(val-target)<1e-4: best2.append((si,names[i],sj,names[j]))
print("MATCH2", best2[:10])
Run probe 14 search
python3 p14.py
adv*tl -0.3551923930644989 adv*d 0.11678817868232727 tl -1.958479642868042 d 0.36872270703315735 kl_tok 0.5788273215293884 full_rev 0.7822656631469727 full_fwd 0.8293264508247375 ratio*adv 0.07552962750196457 adv 0.24970470368862152 TARGET 1.4292851686477661 MATCH1 [] MATCH2 [] [stdout] adv*tl -0.3551923930644989 adv*d 0.11678817868232727 tl -1.958479642868042 d 0.36872270703315735 kl_tok 0.5788273215293884 full_rev 0.7822656631469727 full_fwd 0.8293264508247375 ratio*adv 0.07552962750196457 adv 0.24970470368862152 TARGET 1.4292851686477661 MATCH1 [] MATCH2 []
import torch
import reforacle as R
g = torch.Generator().manual_seed(31)
def rn(*s): return torch.randn(*s, generator=g)
V=5;N=3;T=4
logits=rn(N,T,V); ref=rn(N,T,V); labels=torch.randint(0,V,(N,T),generator=g)
mask=(torch.rand(N,T,generator=g)>0.3).float(); adv=rn(N,T); beta=0.1
lp=torch.log_softmax(logits,-1); rlp=torch.log_softmax(ref,-1)
tl=lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
rtl=rlp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
d=tl-rtl; kl_tok=torch.exp(-d)-1+d
def mm(x): return ((x*mask).sum()/mask.sum()).item()
# isolate: beta=0
print("beta=0:", R.reverse_kl_objective(logits,ref,labels,mask,adv,0.0).item())
print(" mm(-adv*tl)", mm(-adv*tl))
# isolate: adv=0
z=torch.zeros(N,T)
print("adv=0,beta=0.1:", R.reverse_kl_objective(logits,ref,z,mask,z,beta).item())
print("adv=0,beta=1:", R.reverse_kl_objective(logits,ref,labels,mask,z,1.0).item())
print(" mm(kl_tok)", mm(kl_tok), "beta1")
# adv=1 all ones, beta=0
one=torch.ones(N,T)
print("adv=1,beta=0:", R.reverse_kl_objective(logits,ref,labels,mask,one,0.0).item())
print(" mm(-tl)", mm(-tl))
print(" mm(tl)", mm(tl))
Isolate reverse_kl_objective terms
python3 p15.py
Exit code 1
beta=0: 1.2556371688842773
mm(-adv*tl) 0.3551923930644989
Traceback (most recent call last):
File "/app/p15.py", line 19, in <module>
print("adv=0,beta=0.1:", R.reverse_kl_objective(logits,ref,z,mask,z,beta).item())
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "reforacle.py", line 225, in reforacle.reverse_kl_objective
File "reforacle.py", line 38, in reforacle.token_logprobs
RuntimeError: gather(): Expected dtype int64 for index
[error] tool reported failureimport torch
import reforacle as R
g = torch.Generator().manual_seed(31)
def rn(*s): return torch.randn(*s, generator=g)
V=5;N=3;T=4
logits=rn(N,T,V); ref=rn(N,T,V); labels=torch.randint(0,V,(N,T),generator=g)
mask=(torch.rand(N,T,generator=g)>0.3).float(); adv=rn(N,T); beta=0.1
lp=torch.log_softmax(logits,-1); rlp=torch.log_softmax(ref,-1)
tl=lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
rtl=rlp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
d=tl-rtl; kl_tok=torch.exp(-d)-1+d; ratio=torch.exp(d)
def mm(x): return ((x*mask).sum()/mask.sum()).item()
one=torch.ones(N,T); z=torch.zeros(N,T)
F=R.reverse_kl_objective
print("adv=1,beta=0:", F(logits,ref,labels,mask,one,0.0).item(), "| mm(-ratio)",mm(-ratio),"mm(-tl)",mm(-tl))
print("adv=0,beta=0:", F(logits,ref,labels,mask,z,0.0).item())
print("adv=real,beta=0:", F(logits,ref,labels,mask,adv,0.0).item(), "| mm(-ratio*adv)",mm(-ratio*adv))
print("adv=real,beta=1:", F(logits,ref,labels,mask,adv,1.0).item())
print("adv=0,beta=1:", F(logits,ref,labels,mask,z,1.0).item(), "| mm(kl_tok)",mm(kl_tok),"mm(full_rev)", mm((lp.exp()*(lp-rlp)).sum(-1)))
print("adv=0,beta=0.1:", F(logits,ref,labels,mask,z,0.1).item())
Isolate reverse_kl_objective 2
python3 p16.py
adv=1,beta=0: 5.875438213348389 | mm(-ratio) -2.7466092109680176 mm(-tl) 1.958479642868042 adv=0,beta=0: 0.0 adv=real,beta=0: 1.2556371688842773 | mm(-ratio*adv) -0.07552962750196457 adv=real,beta=1: 2.9921188354492188 adv=0,beta=1: 1.736481785774231 | mm(kl_tok) 0.5788273215293884 mm(full_rev) 0.7822656631469727 adv=0,beta=0.1: 0.1736481785774231 [stdout] adv=1,beta=0: 5.875438213348389 | mm(-ratio) -2.7466092109680176 mm(-tl) 1.958479642868042 adv=0,beta=0: 0.0 adv=real,beta=0: 1.2556371688842773 | mm(-ratio*adv) -0.07552962750196457 adv=real,beta=1: 2.9921188354492188 adv=0,beta=1: 1.736481785774231 | mm(kl_tok) 0.5788273215293884 mm(full_rev) 0.7822656631469727 adv=0,beta=0.1: 0.1736481785774231
import torch
import reforacle as R
g = torch.Generator().manual_seed(31)
def rn(*s): return torch.randn(*s, generator=g)
V=5;N=3;T=4
logits=rn(N,T,V); ref=rn(N,T,V); labels=torch.randint(0,V,(N,T),generator=g)
mask=(torch.rand(N,T,generator=g)>0.3).float(); adv=rn(N,T)
lp=torch.log_softmax(logits,-1); rlp=torch.log_softmax(ref,-1)
tl=lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
rtl=rlp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
d=tl-rtl; ratio=torch.exp(d)
def mm(x): return ((x*mask).sum()/mask.sum()).item()
# want mm(w)=5.8754 (adv=1,beta=0 pg part) and mm(adv*w)=1.2556
target1=5.875438213348389; target2=1.2556371688842773
cands={
'-tl':-tl,'-rtl':-rtl,'ratio':ratio,'-ratio':-ratio,'-d':-d,'d':d,
'ratio*-tl':ratio*-tl,'-tl*ratio':-tl*ratio,'-rtl*ratio':-rtl*ratio,
'ratio^2':ratio**2,'-tl-rtl':-tl-rtl,'exp(-tl)':torch.exp(-tl),
'-tl+ratio-1':-tl+ratio-1,'ratio*-d':ratio*-d,
}
for k,v in cands.items():
a=mm(v); b=mm(adv*v)
flag=""
if abs(a-target1)<1e-3: flag+=" T1!"
if abs(b-target2)<1e-3: flag+=" T2!"
print(f"{k:15s} mm={a:.4f} mm(adv*)={b:.4f}{flag}")
Search pg token expression
python3 p17.py
-tl mm=1.9585 mm(adv*)=0.3552 -rtl mm=2.3272 mm(adv*)=0.4720 ratio mm=2.7466 mm(adv*)=0.0755 -ratio mm=-2.7466 mm(adv*)=-0.0755 -d mm=-0.3687 mm(adv*)=-0.1168 d mm=0.3687 mm(adv*)=0.1168 ratio*-tl mm=4.0159 mm(adv*)=0.2487 -tl*ratio mm=4.0159 mm(adv*)=0.2487 -rtl*ratio mm=8.3145 mm(adv*)=-0.4881 ratio^2 mm=17.5304 mm(adv*)=-3.3516 -tl-rtl mm=4.2857 mm(adv*)=0.8272 exp(-tl) mm=10.4459 mm(adv*)=0.4865 -tl+ratio-1 mm=3.7051 mm(adv*)=0.1810 ratio*-d mm=-4.2986 mm(adv*)=0.7368 [stdout] -tl mm=1.9585 mm(adv*)=0.3552 -rtl mm=2.3272 mm(adv*)=0.4720 ratio mm=2.7466 mm(adv*)=0.0755 -ratio mm=-2.7466 mm(adv*)=-0.0755 -d mm=-0.3687 mm(adv*)=-0.1168 d mm=0.3687 mm(adv*)=0.1168 ratio*-tl mm=4.0159 mm(adv*)=0.2487 -tl*ratio mm=4.0159 mm(adv*)=0.2487 -rtl*ratio mm=8.3145 mm(adv*)=-0.4881 ratio^2 mm=17.5304 mm(adv*)=-3.3516 -tl-rtl mm=4.2857 mm(adv*)=0.8272 exp(-tl) mm=10.4459 mm(adv*)=0.4865 -tl+ratio-1 mm=3.7051 mm(adv*)=0.1810 ratio*-d mm=-4.2986 mm(adv*)=0.7368
import torch
import reforacle as R
g = torch.Generator().manual_seed(31)
def rn(*s): return torch.randn(*s, generator=g)
V=5;N=3;T=4
logits=rn(N,T,V); ref=rn(N,T,V); labels=torch.randint(0,V,(N,T),generator=g)
mask=(torch.rand(N,T,generator=g)>0.3).float(); adv=rn(N,T)
lp=torch.log_softmax(logits,-1); rlp=torch.log_softmax(ref,-1)
tl=lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
rtl=rlp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
d=tl-rtl; ratio=torch.exp(d)
S=mask.sum().item()
F=R.reverse_kl_objective
# recover w per token
w=torch.zeros(N,T)
for i in range(N):
for j in range(T):
e=torch.zeros(N,T); e[i,j]=1.0
p=F(logits,ref,labels,mask,e,0.0).item()
if mask[i,j]>0: w[i,j]=p*S/1.0
print("recovered w:\n", w)
print("tl:\n", tl)
print("rtl:\n", rtl)
print("ratio:\n", ratio)
print("-tl:\n", -tl)
# compare w to candidates
print("w / -tl:\n", (w/(-tl)))
print("w vs ratio*-tl:\n", ratio*-tl)
print("w vs -tl:\n", -tl)
Recover per-token weight
python3 p18.py
recovered w:
tensor([[6.0507, 0.0000, 4.2397, 6.5983],
[6.0507, 4.6546, 4.2397, 6.5983],
[6.0507, 4.6546, 4.2397, 6.5983]])
tl:
tensor([[-1.7104, -0.9765, -1.5913, -1.6009],
[-2.8469, -3.4200, -1.4243, -3.4521],
[-1.4933, -1.2346, -1.2241, -1.5452]])
rtl:
tensor([[-1.7693, -1.2353, -2.8993, -1.4342],
[-1.9638, -1.8922, -3.5205, -3.0924],
[-1.4253, -2.3620, -3.5196, -1.7205]])
ratio:
tensor([[1.0606, 1.2955, 3.6988, 0.8464],
[0.4135, 0.2170, 8.1352, 0.6979],
[0.9342, 3.0878, 9.9297, 1.1916]])
-tl:
tensor([[1.7104, 0.9765, 1.5913, 1.6009],
[2.8469, 3.4200, 1.4243, 3.4521],
[1.4933, 1.2346, 1.2241, 1.5452]])
w / -tl:
tensor([[3.5375, 0.0000, 2.6643, 4.1215],
[2.1253, 1.3610, 2.9766, 1.9114],
[4.0518, 3.7702, 3.4635, 4.2701]])
w vs ratio*-tl:
tensor([[ 1.8142, 1.2650, 5.8858, 1.3550],
[ 1.1772, 0.7422, 11.5873, 2.4093],
[ 1.3951, 3.8121, 12.1550, 1.8413]])
w vs -tl:
tensor([[1.7104, 0.9765, 1.5913, 1.6009],
[2.8469, 3.4200, 1.4243, 3.4521],
[1.4933, 1.2346, 1.2241, 1.5452]])
[stdout]
recovered w:
tensor([[6.0507, 0.0000, 4.2397, 6.5983],
[6.0507, 4.6546, 4.2397, 6.5983],
[6.0507, 4.6546, 4.2397, 6.5983]])
tl:
tensor([[-1.7104, -0.9765, -1.5913, -1.6009],
[-2.8469, -3.4200, -1.4243, -3.4521],
[-1.4933, -1.2346, -1.2241, -1.5452]])
rtl:
tensor([[-1.7693, -1.2353, -2.8993, -1.4342],
[-1.9638, -1.8922, -3.5205, -3.0924],
[-1.4253, -2.3620, -3.5196, -1.7205]])
ratio:
tensor([[1.0606, 1.2955, 3.6988, 0.8464],
[0.4135, 0.2170, 8.1352, 0.6979],
[0.9342, 3.0878, 9.9297, 1.1916]])
-tl:
tensor([[1.7104, 0.9765, 1.5913, 1.6009],
[2.8469, 3.4200, 1.4243, 3.4521],
[1.4933, 1.2346, 1.2241, 1.5452]])
w / -tl:
tensor([[3.5375, 0.0000, 2.6643, 4.1215],
[2.1253, 1.3610, 2.9766, 1.9114],
[4.0518, 3.7702, 3.4635, 4.2701]])
w vs ratio*-tl:
tensor([[ 1.8142, 1.2650, 5.8858, 1.3550],
[ 1.1772, 0.7422, 11.5873, 2.4093],
[ 1.3951, 3.8121, 12.1550, 1.8413]])
w vs -tl:
tensor([[1.7104, 0.9765, 1.5913, 1.6009],
[2.8469, 3.4200, 1.4243, 3.4521],
[1.4933, 1.2346, 1.2241, 1.5452]])import torch
import reforacle as R
g = torch.Generator().manual_seed(31)
def rn(*s): return torch.randn(*s, generator=g)
V=5;N=3;T=4
logits=rn(N,T,V); ref=rn(N,T,V); labels=torch.randint(0,V,(N,T),generator=g)
mask=(torch.rand(N,T,generator=g)>0.3).float(); advT=rn(N,T)
lp=torch.log_softmax(logits,-1); rlp=torch.log_softmax(ref,-1)
tl=lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
rtl=rlp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
d=tl-rtl
def mm(x): return ((x*mask).sum()/mask.sum()).item()
F=R.reverse_kl_objective
# cross formula test with advT
S=mask.sum()
cross = -(1/S)*sum( advT[:,j].sum()* (mask[:,j]*tl[:,j]).sum() for j in range(T))
print("cross(advT)", cross.item(), "R(advT,beta0)", F(logits,ref,labels,mask,advT,0.0).item())
# try adv shape (N,)
advN=rn(N)
print("R(advN,beta0)", F(logits,ref,labels,mask,advN,0.0).item())
print(" -mm(advN[:,None]*tl)", mm(-advN[:,None]*tl))
print(" -(advN[:,None]*tl*mask).sum()/N", (-(advN[:,None]*tl*mask).sum()/N).item())
# KL term value with these tensors
print("R(advN,beta1)-R(advN,beta0)", F(logits,ref,labels,mask,advN,1.0).item()-F(logits,ref,labels,mask,advN,0.0).item())
full_rev=(lp.exp()*(lp-rlp)).sum(-1)
full_fwd=(rlp.exp()*(rlp-lp)).sum(-1)
kl_tok=torch.exp(-d)-1+d
print(" mm(full_rev)",mm(full_rev),"mm(full_fwd)",mm(full_fwd),"mm(kl_tok)",mm(kl_tok))
print(" ms(full_rev)/N", (full_rev*mask.any(-1)).sum().item() if False else (full_rev).sum().item()/N)
Test cross formula and adv (N,)
python3 p19.py
cross(advT) 1.2556370496749878 R(advT,beta0) 1.2556371688842773 R(advN,beta0) 0.22671037912368774 -mm(advN[:,None]*tl) 0.22671037912368774 -(advN[:,None]*tl*mask).sum()/N 0.8312714099884033 R(advN,beta1)-R(advN,beta0) 0.5788272619247437 mm(full_rev) 0.7822656631469727 mm(full_fwd) 0.8293264508247375 mm(kl_tok) 0.5788273215293884 ms(full_rev)/N 3.0077622731526694 [stdout] cross(advT) 1.2556370496749878 R(advT,beta0) 1.2556371688842773 R(advN,beta0) 0.22671037912368774 -mm(advN[:,None]*tl) 0.22671037912368774 -(advN[:,None]*tl*mask).sum()/N 0.8312714099884033 R(advN,beta1)-R(advN,beta0) 0.5788272619247437 mm(full_rev) 0.7822656631469727 mm(full_fwd) 0.8293264508247375 mm(kl_tok) 0.5788273215293884 ms(full_rev)/N 3.0077622731526694
import torch
import reforacle as R
g = torch.Generator().manual_seed(41)
def rn(*s): return torch.randn(*s, generator=g)
V=5;N=4;T=3;G=2
logits=rn(N,T,V); old=rn(N,T,V); ref=rn(N,T,V)
labels=torch.randint(0,V,(N,T),generator=g)
cm=(torch.rand(N,T,generator=g)>0.3).float(); rewards=rn(N)
beta=0.1; cl_lo=0.2; cl_hi=0.2
def tlp(lg):
lp=torch.log_softmax(lg,-1); return lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
logp=tlp(logits); oldlp=tlp(old); reflp=tlp(ref)
def mm(x): return (x*cm).sum()/cm.sum()
# advantages
adv=R.grpo_advantages(rewards,G,True) # (N,)
A=adv.unsqueeze(-1)
ratio=torch.exp(logp-oldlp)
l1=ratio*A; l2=torch.clamp(ratio,1-cl_lo,1+cl_hi)*A
pg=-torch.min(l1,l2)
d=logp-reflp; kl_k3=torch.exp(-d)-1+d
mine=mm(pg+beta*kl_k3)
print("my grpo k3", mine.item())
for est in ['k1','k2','k3']:
kl={'k1':d,'k2':0.5*d*d,'k3':torch.exp(-d)-1+d}[est]
print(" my", est, mm(pg+beta*kl).item())
print("R grpo k1", R.grpo_objective(logits,old,ref,labels,cm,rewards,G,beta,cl_lo,cl_hi,True,'k1').item())
print("R grpo k2", R.grpo_objective(logits,old,ref,labels,cm,rewards,G,beta,cl_lo,cl_hi,True,'k2').item())
print("R grpo k3", R.grpo_objective(logits,old,ref,labels,cm,rewards,G,beta,cl_lo,cl_hi,True,'k3').item())
Probe grpo_objective
python3 p20.py
my grpo k3 0.32634857296943665 my k1 0.24630212783813477 my k2 0.31633320450782776 my k3 0.32634857296943665 R grpo k1 0.24630212783813477 R grpo k2 0.31633320450782776 R grpo k3 0.32634857296943665 [stdout] my grpo k3 0.32634857296943665 my k1 0.24630212783813477 my k2 0.31633320450782776 my k3 0.32634857296943665 R grpo k1 0.24630212783813477 R grpo k2 0.31633320450782776 R grpo k3 0.32634857296943665
import torch
import reforacle as R
g = torch.Generator().manual_seed(43)
def rn(*s): return torch.randn(*s, generator=g)
T=5
rewards=rn(T); values=rn(T); old_values=rn(T); logp=rn(T); old_logp=rn(T)
nv=rn(1).item()
gamma,lam=0.99,0.95; cl_lo,cl_hi=0.2,0.2; vf_clip=0.2; vf_coef=0.5
Rv=R.ppo_objective(rewards,values,old_values,logp,old_logp,nv,gamma,lam,cl_lo,cl_hi,vf_clip,vf_coef)
print("R ppo", Rv.item() if hasattr(Rv,'item') else Rv)
# my
adv=R.gae(rewards,values,nv,gamma,lam)
returns=adv+values
ratio=torch.exp(logp-old_logp)
def build(adv_use):
l1=ratio*adv_use; l2=torch.clamp(ratio,1-cl_lo,1+cl_hi)*adv_use
pg=-torch.min(l1,l2)
return pg.mean()
# value loss
vc=old_values+torch.clamp(values-old_values,-vf_clip,vf_clip)
vl=0.5*torch.max((values-returns)**2,(vc-returns)**2).mean()
for name,advu in [("raw",adv),("whiten",(adv-adv.mean())/(adv.std(unbiased=True)+1e-8)),("whitenF",(adv-adv.mean())/(adv.std(unbiased=False)+1e-8))]:
total=build(advu)+vf_coef*vl
print(f" my adv={name}: pg={build(advu).item():.6f} total={total.item():.6f}")
print(" vl", vl.item())
Probe ppo_objective
python3 p21.py
R ppo 0.9653177261352539 my adv=raw: pg=-1.647207 total=-0.551150 my adv=whiten: pg=-0.116937 total=0.979120 my adv=whitenF: pg=-0.130739 total=0.965318 vl 2.1921145915985107 [stdout] R ppo 0.9653177261352539 my adv=raw: pg=-1.647207 total=-0.551150 my adv=whiten: pg=-0.116937 total=0.979120 my adv=whitenF: pg=-0.130739 total=0.965318 vl 2.1921145915985107
from __future__ import annotations
import math
import torch
import torch.nn.functional as F
def masked_mean(values, mask, axis=None):
mask = mask.to(values.dtype)
num = (values * mask).sum(axis)
den = mask.sum(axis)
return num / den.clamp_min(1e-12)
def masked_sum(values, mask, axis=None):
mask = mask.to(values.dtype)
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):
lp = torch.log_softmax(logits, -1)
return lp.gather(-1, labels.long().unsqueeze(-1)).squeeze(-1)
def selective_logprobs(logits, labels, mask):
tl = token_logprobs(logits, labels)
return (tl * mask.to(tl.dtype)).sum(-1)
def sequence_logprob(logits, labels, mask, length_normalize):
tl = token_logprobs(logits, labels)
mask = mask.to(tl.dtype)
s = (tl * mask).sum(-1)
if length_normalize:
return s / mask.sum(-1).clamp_min(1e-12)
return s
def entropy(logits, mask):
lp = torch.log_softmax(logits, -1)
ent = -(lp.exp() * lp).sum(-1)
return masked_mean(ent, mask)
def dpo_loss(pc, pr, rc, rr, beta, label_smoothing):
logits = beta * ((pc - rc) - (pr - rr))
loss = -F.logsigmoid(logits) * (1 - label_smoothing) - F.logsigmoid(-logits) * label_smoothing
return loss.mean()
def ipo_loss(pc, pr, rc, rr, beta):
logits = (pc - rc) - (pr - rr)
return ((logits - 1.0 / (2.0 * beta)) ** 2).mean()
def grpo_advantages(rewards, group_size, scale_by_std):
g = rewards.view(-1, group_size)
mean = g.mean(-1, keepdim=True)
centered = g - mean
if scale_by_std:
std = g.std(-1, unbiased=False, keepdim=True)
out = centered / (std + 1e-6)
else:
out = centered
return out.reshape(rewards.shape)
def gae(rewards, values, next_value, gamma, lam):
T = rewards.shape[0]
adv = torch.zeros_like(rewards)
last = torch.zeros((), dtype=rewards.dtype)
vnext = next_value if torch.is_tensor(next_value) else torch.tensor(next_value, dtype=rewards.dtype)
for t in range(T - 1, -1, -1):
delta = rewards[t] + gamma * vnext - values[t]
last = delta + gamma * lam * last
adv[t] = last
vnext = values[t]
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 + d
raise ValueError(estimator)
def clipped_pg_loss(logp, old_logp, advantages, mask, clip_low, clip_high):
ratio = torch.exp(logp - old_logp)
l1 = ratio * advantages
l2 = torch.clamp(ratio, 1 - clip_low, 1 + clip_high) * advantages
loss = -torch.min(l1, l2)
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(values, mask, shift_mean):
mask = mask.to(values.dtype)
mean = masked_mean(values, mask)
var = masked_mean((values - mean) ** 2, mask)
out = (values - mean) / torch.sqrt(var + 1e-8)
if not shift_mean:
out = out + mean
return out
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):
logp = token_logprobs(logits, labels)
old_logp = token_logprobs(old_logits, labels)
ref_logp = token_logprobs(ref_logits, labels)
adv = grpo_advantages(rewards, group_size, scale_by_std).unsqueeze(-1)
ratio = torch.exp(logp - old_logp)
l1 = ratio * adv
l2 = torch.clamp(ratio, 1 - clip_low, 1 + clip_high) * adv
pg = -torch.min(l1, l2)
kl = kl_penalty(logp, ref_logp, kl_estimator)
return masked_mean(pg + beta * kl, completion_mask)
def ppo_objective(rewards, values, old_values, logp, old_logp, next_value,
gamma, lam, clip_low, clip_high, vf_clip, vf_coef):
adv = gae(rewards, values, next_value, gamma, lam)
returns = adv + values
mean = adv.mean()
std = adv.std(unbiased=False)
adv_w = (adv - mean) / (std + 1e-8)
ratio = torch.exp(logp - old_logp)
l1 = ratio * adv_w
l2 = torch.clamp(ratio, 1 - clip_low, 1 + clip_high) * adv_w
pg = -torch.min(l1, l2).mean()
vl = value_loss(values, old_values, returns, vf_clip)
return pg + vf_coef * vl
def rloo_advantages(rewards, group_size):
g = rewards.view(-1, group_size)
total = g.sum(-1, keepdim=True)
baseline = (total - g) / (group_size - 1)
return (g - baseline).reshape(rewards.shape)
def reverse_kl(logp, ref_logp):
diff = ref_logp - logp
return torch.exp(diff) - diff - 1
def importance_ratio(logp, old_logp, clip):
ratio = torch.exp(logp - old_logp)
if clip is not None:
ratio = torch.clamp(ratio, 1 - clip, 1 + clip)
return ratio
def discounted_returns(rewards, gamma):
T = rewards.shape[0]
out = torch.zeros_like(rewards)
acc = torch.zeros((), dtype=rewards.dtype)
for t in range(T - 1, -1, -1):
acc = rewards[t] + gamma * acc
out[t] = acc
return out
def normalize(x, eps):
return (x - x.mean()) / (x.std() + eps)
def top_p_mask(probs, p):
sorted_probs, sorted_idx = torch.sort(probs, dim=-1, descending=True)
cumsum = sorted_probs.cumsum(-1)
remove = cumsum - 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(-1, keepdim=True)
return out
def smoothed_nll(logits, labels, smoothing):
lp = torch.log_softmax(logits, -1)
nll = -lp.gather(-1, labels.long().unsqueeze(-1)).squeeze(-1)
smooth = -lp.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):
logp = token_logprobs(logits, labels)
old_logp = token_logprobs(old_logits, labels)
adv = rloo_advantages(rewards, group_size).unsqueeze(-1)
ratio = torch.exp(logp - old_logp)
l1 = ratio * adv
l2 = torch.clamp(ratio, 1 - clip_low, 1 + clip_high) * adv
loss = -torch.min(l1, l2)
return masked_mean(loss, mask)
def reverse_kl_objective(logits, ref_logits, labels, mask, advantages, beta):
logp = token_logprobs(logits, labels)
ref_logp = token_logprobs(ref_logits, labels)
adv = advantages
if adv.dim() == 1:
adv = adv.unsqueeze(-1)
d = logp - ref_logp
kl = torch.exp(-d) - 1 + d
pg = -adv * logp
return masked_mean(pg + beta * kl, mask)
def cross_entropy(logits, labels, ignore_index):
V = logits.shape[-1]
return F.cross_entropy(logits.reshape(-1, V), labels.reshape(-1).long(),
ignore_index=ignore_index)
def top_k_mask(logits, k):
kth = torch.topk(logits, k, dim=-1).values[..., -1:]
return logits >= kth
def group_mean_baseline(rewards, group_size):
g = rewards.view(-1, group_size)
mean = g.mean(-1, keepdim=True)
return (g - mean).reshape(rewards.shape)
def lambda_returns(rewards, values, next_value, gamma, lam):
return gae(rewards, values, next_value, gamma, lam) + values
def symmetric_kl(logp, ref_logp):
return torch.cosh(logp - ref_logp) - 1
def huber_value_loss(values, returns, delta):
return F.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).abs() > clip).float().mean()
def masked_whiten(values, mask, shift_mean):
return whiten(values, mask, shift_mean)
def logprob_at_temperature(logits, labels, temperature):
lp = torch.log_softmax(logits / temperature, -1)
return lp.gather(-1, labels.long().unsqueeze(-1)).squeeze(-1)
def advantage_mean_std(advantages, mask):
mask = mask.to(advantages.dtype)
mean = masked_mean(advantages, mask)
var = masked_mean((advantages - mean) ** 2, mask)
return torch.stack([mean, torch.sqrt(var)])
def argmax_tokens(logits):
return logits.argmax(-1)
def mode_label(labels):
return torch.mode(labels).values
def median_reward(rewards):
return torch.median(rewards)
def masked_max(values, mask):
filled = values.masked_fill(mask == 0, float("-inf"))
return filled.max()
def quantile_lower(x, q):
return torch.quantile(x, q, interpolation="lower")
def pad_mask_from_lengths(lengths, max_len):
ar = torch.arange(max_len, device=lengths.device)
return (ar.unsqueeze(0) < lengths.unsqueeze(-1)).to(torch.float32)
def first_nonzero_index(mask):
nz = mask != 0
has = nz.any(-1)
first = torch.argmax(nz.to(torch.int64), dim=-1)
return torch.where(has, first, torch.full_like(first, -1))
def cumulative_max(x):
return torch.cummax(x, dim=-1).values
def bucketize_reward(rewards, boundaries):
return torch.bucketize(rewards, boundaries)
import torch
import reforacle as R
import objectives as O
g = torch.Generator().manual_seed(0)
def rn(*s): return torch.randn(*s, generator=g)
def rint(hi,*s): return torch.randint(0,hi,s,generator=g)
def rmask(*s): return (torch.rand(*s,generator=g)>0.3).float()
def cmp(name, a, b, atol=1e-5, rtol=1e-4):
a=a if torch.is_tensor(a) else torch.tensor(a)
b=b if torch.is_tensor(b) else torch.tensor(b)
a=a.float(); b=b.float()
if a.shape!=b.shape:
print(f"FAIL {name}: shape {a.shape} vs {b.shape}"); return False
ok=torch.allclose(a,b,atol=atol,rtol=rtol)
if not ok:
print(f"FAIL {name}: max diff {(a-b).abs().max().item():.2e}\n R={a.flatten()[:6]}\n O={b.flatten()[:6]}")
return ok
results={}
def T(name, fn_args_list, **kw):
allok=True
for args in fn_args_list:
try:
ra=getattr(R,name)(*args)
except Exception as e:
print(f"REF ERR {name}: {e}"); allok=False; continue
try:
oa=getattr(O,name)(*args)
except Exception as e:
print(f"IMPL ERR {name}: {e}"); allok=False; continue
if not cmp(name, ra, oa, **kw): allok=False
results[name]=allok
for trial in range(5):
v=rn(3,4); m=rmask(3,4)
T('masked_mean',[(v,m),(v,m,0),(v,m,1),(rn(5),torch.zeros(5))])
T('masked_sum',[(v,m),(v,m,1),(v,m,0)])
x=rn(4,6)
T('logsumexp',[(x,1),(x,0),(x,-1)])
T('log_softmax',[(x,1),(x,-1)])
B,Tt,Vv=2,4,7
lg=rn(B,Tt,Vv); lb=rint(Vv,B,Tt); mk=rmask(B,Tt)
T('token_logprobs',[(lg,lb)])
T('selective_logprobs',[(lg,lb,mk)])
T('sequence_logprob',[(lg,lb,mk,False),(lg,lb,mk,True)])
T('entropy',[(lg,mk)])
T('normalized_entropy',[(lg,mk)])
pc,pr,rc,rr=rn(5),rn(5),rn(5),rn(5)
T('dpo_loss',[(pc,pr,rc,rr,0.1,0.0),(pc,pr,rc,rr,0.5,0.1)])
T('ipo_loss',[(pc,pr,rc,rr,0.3),(pc,pr,rc,rr,1.0)])
rw=rn(6)
T('grpo_advantages',[(rw,3,True),(rw,3,False),(rw,2,True)])
T('rloo_advantages',[(rw,3),(rw,2)])
T('group_mean_baseline',[(rw,3),(rw,2)])
rwT=rn(5); vals=rn(5); nv=rn(1).item()
T('gae',[(rwT,vals,nv,0.99,0.95),(rwT,vals,rn(1),0.9,0.8)])
T('lambda_returns',[(rwT,vals,nv,0.99,0.95)])
T('discounted_returns',[(rwT,0.99),(rwT,0.5)])
lp=rn(3,4); ref=rn(3,4); old=rn(3,4)
T('kl_penalty',[(lp,ref,'k1'),(lp,ref,'k2'),(lp,ref,'k3')])
T('reverse_kl',[(lp,ref)])
T('symmetric_kl',[(lp,ref)])
T('importance_ratio',[(lp,old,None),(lp,old,0.2)])
T('clip_fraction',[(lp,old,0.2),(lp,old,0.5)])
adv=rn(3,4)
T('clipped_pg_loss',[(lp,old,adv,m,0.2,0.2),(lp,old,adv,m,0.1,0.3)])
vv=rn(3,4); ov=rn(3,4); ret=rn(3,4)
T('value_loss',[(vv,ov,ret,0.2)])
T('huber_value_loss',[(rn(6),rn(6),1.0),(rn(6),rn(6),0.5)])
xw=rn(8); mw=rmask(8)
T('whiten',[(xw,mw,True),(xw,mw,False)])
T('masked_whiten',[(xw,mw,True),(xw,mw,False)])
T('normalize',[(rn(7),1e-8)])
T('advantage_mean_std',[(rn(8),rmask(8))])
# dpo_sequence_loss
args=(rn(2,4,6),rn(2,4,6),rn(2,4,6),rn(2,4,6),rint(6,2,4),rint(6,2,4),rmask(2,4),rmask(2,4),0.1,0.0)
T('dpo_sequence_loss',[args])
# grpo_objective
N=4
go=(rn(N,3,6),rn(N,3,6),rn(N,3,6),rint(6,N,3),rmask(N,3),rn(N),2,0.1,0.2,0.2,True,'k3')
T('grpo_objective',[go, (rn(N,3,6),rn(N,3,6),rn(N,3,6),rint(6,N,3),rmask(N,3),rn(N),2,0.05,0.2,0.2,False,'k1')])
# ppo_objective
po=(rn(5),rn(5),rn(5),rn(5),rn(5),rn(1).item(),0.99,0.95,0.2,0.2,0.2,0.5)
T('ppo_objective',[po])
# rloo_objective
ro=(rn(N,3,6),rn(N,3,6),rint(6,N,3),rmask(N,3),rn(N),2,0.2,0.2)
T('rloo_objective',[ro])
# reverse_kl_objective (adv (N,))
rk=(rn(N,3,6),rn(N,3,6),rint(6,N,3),rmask(N,3),rn(N),0.1)
T('reverse_kl_objective',[rk])
T('bradley_terry_logit',[(torch.tensor(1.0),torch.tensor(0.3),0.5),(rn(4),rn(4),0.7)])
pr2=torch.softmax(rn(3,8),-1)
T('top_p_mask',[(pr2,0.8),(pr2,0.5),(pr2,0.95)])
T('top_k_mask',[(rn(3,8),3),(rn(3,8),1)])
T('smoothed_nll',[(rn(2,4,7),rint(7,2,4),0.1)])
T('logprob_at_temperature',[(rn(2,4,7),rint(7,2,4),2.0),(rn(2,4,7),rint(7,2,4),0.5)])
T('cross_entropy',[(rn(2,4,7),rint(7,2,4),-100),(rn(5,7),rint(7,5),-100)])
T('argmax_tokens',[(rn(2,4,7),)])
T('mode_label',[(rint(5,9),)])
T('median_reward',[(rn(7),),(rn(8),)])
T('masked_max',[(rn(8),rmask(8))])
T('quantile_lower',[(rn(10),0.25),(rn(10),0.5),(rn(10),0.75)])
T('pad_mask_from_lengths',[(torch.tensor([2,4,1,5]),6)])
T('first_nonzero_index',[(torch.tensor([[0.,0,1,0],[0,0,0,0],[1,0,0,0]]),)])
T('cumulative_max',[(rn(6),),(rn(3,5),)])
T('bucketize_reward',[(rn(6),torch.tensor([-0.5,0.,0.5,1.0]))])
fails=[k for k,v in results.items() if not v]
print("\n==== SUMMARY ====")
print(f"tested {len(results)} functions, {len(fails)} failing")
if fails: print("FAILING:", fails)
else: print("ALL PASS")
Run validation
python3 validate.py 2>&1 | tail -40
FAIL cumulative_max: max diff 1.14e+00 R=tensor([-0.9800, -0.9800, -0.0569, 1.0851, 1.0851, 1.0851]) O=tensor([-0.9800, -0.0569, 1.0851, 1.0851, 1.0851, 1.0851]) FAIL cumulative_max: max diff 1.32e+00 R=tensor([-0.5246, -0.4470, 0.3584, 0.7387, 0.3398, -0.5246]) O=tensor([-0.5246, -0.4470, 0.3584, 0.7387, 0.7387, 0.6358]) FAIL masked_whiten: max diff 8.04e-02 R=tensor([-1.6969, -0.0000, 0.5808, 0.7561, 1.1085, -1.2456]) O=tensor([-1.6969, -0.0804, 0.5808, 0.7561, 1.1085, -1.2456]) FAIL masked_whiten: max diff 8.04e-02 R=tensor([-1.4204, 0.2765, 0.8573, 1.0326, 1.3850, -0.9691]) O=tensor([-1.4204, 0.1961, 0.8573, 1.0326, 1.3850, -0.9691]) FAIL median_reward: max diff 4.50e-02 R=tensor([0.1988]) O=tensor([0.1538]) FAIL cumulative_max: max diff 1.39e+00 R=tensor([-1.0628, -1.0628, -0.1223, -0.1223, -0.1223, 0.1445]) O=tensor([-1.0628, -0.1223, -0.1223, -0.1223, 0.1445, 1.5361]) FAIL cumulative_max: max diff 3.41e+00 R=tensor([-0.7751, 1.5088, -1.9000, 1.0851, -0.2575, -0.7751]) O=tensor([-0.7751, 1.5088, 1.5088, 1.5088, 1.5088, -0.5414]) FAIL masked_whiten: max diff 6.20e-01 R=tensor([ 0.0000, 0.2433, -1.1602, 0.5786, -1.0963, 1.4346]) O=tensor([ 0.2693, 0.2433, -1.1602, 0.5786, -1.0963, 1.4346]) FAIL masked_whiten: max diff 6.20e-01 R=tensor([ 0.5340, 0.7773, -0.6262, 1.1125, -0.5623, 1.9686]) O=tensor([ 0.8033, 0.7773, -0.6262, 1.1125, -0.5623, 1.9686]) FAIL median_reward: max diff 6.06e-03 R=tensor([-0.3448]) O=tensor([-0.3509]) FAIL cumulative_max: max diff 1.81e+00 R=tensor([-2.9070, -2.9070, -1.1159, 0.6969, 0.6969, 0.6969]) O=tensor([-2.9070, -1.1159, 0.6969, 0.6969, 0.6969, 0.6969]) FAIL cumulative_max: max diff 2.34e+00 R=tensor([-0.8676, -1.1647, -1.1309, 0.9099, 1.0260, -0.8676]) O=tensor([-0.8676, -0.8676, -0.8676, 0.9099, 1.0260, 0.1840]) ==== SUMMARY ==== tested 49 functions, 3 failing FAILING: ['masked_whiten', 'median_reward', 'cumulative_max'] [stdout] FAIL cumulative_max: max diff 1.14e+00 R=tensor([-0.9800, -0.9800, -0.0569, 1.0851, 1.0851, 1.0851]) O=tensor([-0.9800, -0.0569, 1.0851, 1.0851, 1.0851, 1.0851]) FAIL cumulative_max: max diff 1.32e+00 R=tensor([-0.5246, -0.4470, 0.3584, 0.7387, 0.3398, -0.5246]) O=tensor([-0.5246, -0.4470, 0.3584, 0.7387, 0.7387, 0.6358]) FAIL masked_whiten: max diff 8.04e-02 R=tensor([-1.6969, -0.0000, 0.5808, 0.7561, 1.1085, -1.2456]) O=tensor([-1.6969, -0.0804, 0.5808, 0.7561, 1.1085, -1.2456]) FAIL masked_whiten: max diff 8.04e-02 R=tensor([-1.4204, 0.2765, 0.8573, 1.0326, 1.3850, -0.9691]) O=tensor([-1.4204, 0.1961, 0.8573, 1.0326, 1.3850, -0.9691]) FAIL median_reward: max diff 4.50e-02 R=tensor([0.1988]) O=tensor([0.1538]) FAIL cumulative_max: max diff 1.39e+00 R=tensor([-1.0628, -1.0628, -0.1223, -0.1223, -0.1223, 0.1445]) O=tensor([-1.0628, -0.1223, -0.1223, -0.1223, 0.1445, 1.5361]) FAIL cumulative_max: max diff 3.41e+00 R=tensor([-0.7751, 1.5088, -1.9000, 1.0851, -0.2575, -0.7751]) O=tensor([-0.7751, 1.5088, 1.5088, 1.5088, 1.5088, -0.5414]) FAIL masked_whiten: max diff 6.20e-01 R=tensor([ 0.0000, 0.2433, -1.1602, 0.5786, -1.0963, 1.4346]) O=tensor([ 0.2693, 0.2433, -1.1602, 0.5786, -1.0963, 1.4346]) FAIL masked_whiten: max diff 6.20e-01 R=tensor([ 0.5340, 0.7773, -0.6262, 1.1125, -0.5623, 1.9686]) O=tensor([ 0.8033, 0.7773, -0.6262, 1.1125, -0.5623, 1.9686]) FAIL median_reward: max diff 6.06e-03 R=tensor([-0.3448]) O=tensor([-0.3509]) FAIL cumulative_max: max diff 1.81e+00 R=tensor([-2.9070, -2.9070, -1.1159, 0.6969, 0.6969, 0.6969]) O=tensor([-2.9070, -1.1159, 0.6969, 0.6969, 0.6969, 0.6969]) FAIL cumulative_max: max diff 2.34e+00 R=tensor([-0.8676, -1.1647, -1.1309, 0.9099, 1.0260, -0.8676]) O=tensor([-0.8676, -0.8676, -0.8676, 0.9099, 1.0260, 0.1840]) ==== SUMMARY ==== tested 49 functions, 3 failing FAILING: ['masked_whiten', 'median_reward', 'cumulative_max']
import torch
import reforacle as R
x=torch.tensor([1.,3,2,5,4])
print("cummax 1d", R.cumulative_max(x).tolist())
x2=torch.tensor([[1.,3,2,5,4],[5.,1,4,2,3]])
print("cummax 2d\n", R.cumulative_max(x2).tolist())
# hypotheses
c=torch.cummax(x2,dim=-1).values
print("inclusive dim-1\n", c.tolist())
c0=torch.cummax(x2,dim=0).values
print("inclusive dim0\n", c0.tolist())
# exclusive dim -1
ex=torch.empty_like(x2); ex[...,0]=x2[...,0]; ex[...,1:]=c[...,:-1]
print("exclusive dim-1\n", ex.tolist())
# median even
for n in [7,8,6,4]:
g=torch.Generator().manual_seed(n)
r=torch.randn(n,generator=g)
print(f"n={n} R.median={R.median_reward(r).item():.6f} torch.median={torch.median(r).item():.6f} quantile.5={torch.quantile(r,0.5).item():.6f}")
Probe cumulative_max and median
python3 p22.py
cummax 1d [1.0, 1.0, 3.0, 3.0, 5.0] cummax 2d [[1.0, 3.0, 2.0, 5.0, 4.0], [1.0, 3.0, 2.0, 5.0, 4.0]] inclusive dim-1 [[1.0, 3.0, 3.0, 5.0, 5.0], [5.0, 5.0, 5.0, 5.0, 5.0]] inclusive dim0 [[1.0, 3.0, 2.0, 5.0, 4.0], [5.0, 3.0, 4.0, 5.0, 4.0]] exclusive dim-1 [[1.0, 1.0, 3.0, 3.0, 5.0], [5.0, 5.0, 5.0, 5.0, 5.0]] n=7 R.median=-0.146795 torch.median=-0.146795 quantile.5=-0.146795 n=8 R.median=-0.123815 torch.median=-0.143314 quantile.5=-0.123815 n=6 R.median=-0.846059 torch.median=-0.993663 quantile.5=-0.846059 n=4 R.median=0.539890 torch.median=0.232486 quantile.5=0.539890 [stdout] cummax 1d [1.0, 1.0, 3.0, 3.0, 5.0] cummax 2d [[1.0, 3.0, 2.0, 5.0, 4.0], [1.0, 3.0, 2.0, 5.0, 4.0]] inclusive dim-1 [[1.0, 3.0, 3.0, 5.0, 5.0], [5.0, 5.0, 5.0, 5.0, 5.0]] inclusive dim0 [[1.0, 3.0, 2.0, 5.0, 4.0], [5.0, 3.0, 4.0, 5.0, 4.0]] exclusive dim-1 [[1.0, 1.0, 3.0, 3.0, 5.0], [5.0, 5.0, 5.0, 5.0, 5.0]] n=7 R.median=-0.146795 torch.median=-0.146795 quantile.5=-0.146795 n=8 R.median=-0.123815 torch.median=-0.143314 quantile.5=-0.123815 n=6 R.median=-0.846059 torch.median=-0.993663 quantile.5=-0.846059 n=4 R.median=0.539890 torch.median=0.232486 quantile.5=0.539890
/app/objectives.py
def masked_whiten(values, mask, shift_mean):
return whiten(values, mask, shift_mean)def masked_whiten(values, mask, shift_mean):
return whiten(values, mask, shift_mean) * mask.to(values.dtype)/app/objectives.py
def median_reward(rewards):
return torch.median(rewards)def median_reward(rewards):
return torch.quantile(rewards, 0.5)/app/objectives.py
def cumulative_max(x):
return torch.cummax(x, dim=-1).valuesdef cumulative_max(x):
c = torch.cummax(x, dim=0).values
out = torch.empty_like(x)
out[0] = x[0]
out[1:] = c[:-1]
return outRe-run validation
python3 validate.py 2>&1 | tail -15
R=tensor([-1.6019, 1.5838, -0.0371, 0.7350, 0.2713, 0.0404]) O=tensor([-1.6019, 1.5838, -0.0371, 0.7350, -0.0000, 0.0404]) FAIL masked_whiten: max diff 3.76e-01 R=tensor([ 1.3046, -0.0587, -0.3917, -1.9614, 0.4764, -0.8335]) O=tensor([ 1.3046, -0.0587, -0.3917, -1.9614, 0.4764, -0.8335]) FAIL masked_whiten: max diff 2.76e-01 R=tensor([-1.4204, 0.2765, 0.8573, 1.0326, 1.3850, -0.9691]) O=tensor([-1.4204, 0.0000, 0.8573, 1.0326, 1.3850, -0.9691]) FAIL masked_whiten: max diff 5.34e-01 R=tensor([ 0.5340, 0.7773, -0.6262, 1.1125, -0.5623, 1.9686]) O=tensor([ 0.0000, 0.7773, -0.6262, 1.1125, -0.5623, 1.9686]) ==== SUMMARY ==== tested 49 functions, 1 failing FAILING: ['masked_whiten'] [stdout] R=tensor([-1.6019, 1.5838, -0.0371, 0.7350, 0.2713, 0.0404]) O=tensor([-1.6019, 1.5838, -0.0371, 0.7350, -0.0000, 0.0404]) FAIL masked_whiten: max diff 3.76e-01 R=tensor([ 1.3046, -0.0587, -0.3917, -1.9614, 0.4764, -0.8335]) O=tensor([ 1.3046, -0.0587, -0.3917, -1.9614, 0.4764, -0.8335]) FAIL masked_whiten: max diff 2.76e-01 R=tensor([-1.4204, 0.2765, 0.8573, 1.0326, 1.3850, -0.9691]) O=tensor([-1.4204, 0.0000, 0.8573, 1.0326, 1.3850, -0.9691]) FAIL masked_whiten: max diff 5.34e-01 R=tensor([ 0.5340, 0.7773, -0.6262, 1.1125, -0.5623, 1.9686]) O=tensor([ 0.0000, 0.7773, -0.6262, 1.1125, -0.5623, 1.9686]) ==== SUMMARY ==== tested 49 functions, 1 failing FAILING: ['masked_whiten']
import torch
import reforacle as R
g=torch.Generator().manual_seed(77)
x=torch.randn(8,generator=g)
mask=(torch.rand(8,generator=g)>0.4).float()
print("x ", [round(v,4) for v in x.tolist()])
print("mask", mask.tolist())
mw_T=R.masked_whiten(x,mask,True)
mw_F=R.masked_whiten(x,mask,False)
w_T=R.whiten(x,mask,True)
print("whiten T ", [round(v,4) for v in w_T.tolist()])
print("masked_whiten T", [round(v,4) for v in mw_T.tolist()])
print("masked_whiten F", [round(v,4) for v in mw_F.tolist()])
def mm(v,m): return (v*m).sum()/m.sum()
mean=mm(x,mask)
# candidate: unbiased var (n-1)
n=mask.sum()
var_ub=((x-mean)**2*mask).sum()/(n-1)
var_pop=mm((x-mean)**2,mask)
for nm,var in [("pop",var_pop),("ub",var_ub)]:
cand=(x-mean)/torch.sqrt(var+1e-8)
print(f"cand {nm} shiftT", [round(v,4) for v in cand.tolist()])
Inspect masked_whiten exactly
python3 p23.py
x [-0.3568, 0.6007, -0.6968, -0.5242, 0.9087, -1.6423, 0.4583, -0.1266] mask [1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 0.0, 1.0] whiten T [-0.1203, 1.1007, -0.5538, -0.3337, 1.4933, -1.7595, 0.9191, 0.1733] masked_whiten T [-0.1203, 1.1007, -0.5538, -0.3337, 1.4933, -1.7595, 0.0, 0.1733] masked_whiten F [-0.3828, 0.8382, -0.8163, -0.5962, 1.2309, -2.022, -0.2625, -0.0892] cand pop shiftT [-0.1203, 1.1007, -0.5538, -0.3337, 1.4933, -1.7595, 0.9191, 0.1733] cand ub shiftT [-0.1114, 1.019, -0.5127, -0.309, 1.3826, -1.629, 0.8509, 0.1604] [stdout] x [-0.3568, 0.6007, -0.6968, -0.5242, 0.9087, -1.6423, 0.4583, -0.1266] mask [1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 0.0, 1.0] whiten T [-0.1203, 1.1007, -0.5538, -0.3337, 1.4933, -1.7595, 0.9191, 0.1733] masked_whiten T [-0.1203, 1.1007, -0.5538, -0.3337, 1.4933, -1.7595, 0.0, 0.1733] masked_whiten F [-0.3828, 0.8382, -0.8163, -0.5962, 1.2309, -2.022, -0.2625, -0.0892] cand pop shiftT [-0.1203, 1.1007, -0.5538, -0.3337, 1.4933, -1.7595, 0.9191, 0.1733] cand ub shiftT [-0.1114, 1.019, -0.5127, -0.309, 1.3826, -1.629, 0.8509, 0.1604]
/app/objectives.py
def masked_whiten(values, mask, shift_mean):
return whiten(values, mask, shift_mean) * mask.to(values.dtype)def masked_whiten(values, mask, shift_mean):
mask = mask.to(values.dtype)
mean = masked_mean(values, mask)
var = masked_mean((values - mean) ** 2, mask)
core = (values - mean) / torch.sqrt(var + 1e-8)
out = core * mask
if not shift_mean:
out = out + mean
return outRe-run validation
python3 validate.py 2>&1 | tail -6
==== SUMMARY ==== tested 49 functions, 0 failing ALL PASS [stdout] ==== SUMMARY ==== tested 49 functions, 0 failing ALL PASS
import torch
import reforacle as R
import objectives as O
def eq(name,a,b,atol=1e-5,rtol=1e-4):
a=a if torch.is_tensor(a) else torch.tensor(a)
b=b if torch.is_tensor(b) else torch.tensor(b)
ok = a.shape==b.shape and torch.allclose(a.float(),b.float(),atol=atol,rtol=rtol)
print(("OK " if ok else "FAIL")+f" {name}", "" if ok else f"\n R={a}\n O={b}")
# bucketize boundary (exactly on boundary)
bd=torch.tensor([0.,1.,2.])
rv=torch.tensor([0.,1.,2.,-0.0,0.9999,1.0001])
eq("bucketize boundary", R.bucketize_reward(rv,bd), O.bucketize_reward(rv,bd))
# top_p extremes
p=torch.softmax(torch.randn(2,5,generator=torch.Generator().manual_seed(3)),-1)
eq("top_p p=1.0", R.top_p_mask(p,1.0), O.top_p_mask(p,1.0))
eq("top_p p=0.0", R.top_p_mask(p,0.0), O.top_p_mask(p,0.0))
# first_nonzero all zero
mm=torch.tensor([[0.,0,0],[1.,0,1]])
eq("first_nonzero", R.first_nonzero_index(mm), O.first_nonzero_index(mm))
# masked_mean neg axis + all-zero row
g=torch.Generator().manual_seed(4)
v=torch.randn(3,4,generator=g); mk=torch.tensor([[1.,0,1,0],[0,0,0,0],[1,1,1,1]])
eq("masked_mean axis-1 zerorow", R.masked_mean(v,mk,-1), O.masked_mean(v,mk,-1))
eq("masked_mean axis1", R.masked_mean(v,mk,1), O.masked_mean(v,mk,1))
# quantile various
x=torch.randn(11,generator=g)
for q in [0.1,0.33,0.9,1.0,0.0]:
eq(f"quantile {q}", R.quantile_lower(x,q), O.quantile_lower(x,q))
# mode with tie (lowest?)
lab=torch.tensor([3,1,1,3,2])
eq("mode tie", R.mode_label(lab), O.mode_label(lab))
# cross_entropy 2D and reduction
lg=torch.randn(6,5,generator=g); lb=torch.randint(0,5,(6,),generator=g)
eq("ce 2d", R.cross_entropy(lg,lb,-100), O.cross_entropy(lg,lb,-100))
# importance_ratio clip None
lp=torch.randn(4,generator=g); ol=torch.randn(4,generator=g)
eq("imp None", R.importance_ratio(lp,ol,None), O.importance_ratio(lp,ol,None))
# gae with next_value as tensor
rw=torch.randn(6,generator=g); vl=torch.randn(6,generator=g)
eq("gae tensor nv", R.gae(rw,vl,torch.tensor(0.5),0.97,0.9), O.gae(rw,vl,torch.tensor(0.5),0.97,0.9))
eq("gae float nv", R.gae(rw,vl,0.5,0.97,0.9), O.gae(rw,vl,0.5,0.97,0.9))
# reverse_kl_objective adv 1D check big
N=6
rk=(torch.randn(N,4,7,generator=g),torch.randn(N,4,7,generator=g),torch.randint(0,7,(N,4),generator=g),(torch.rand(N,4,generator=g)>0.3).float(),torch.randn(N,generator=g),0.2)
eq("rev_kl_obj", R.reverse_kl_objective(*rk), O.reverse_kl_objective(*rk))
# entropy empty mask
lg2=torch.randn(2,3,5,generator=g)
eq("entropy allzero", R.entropy(lg2, torch.zeros(2,3)), O.entropy(lg2, torch.zeros(2,3)))
# masked_max 2d
vv=torch.randn(3,4,generator=g); m2=(torch.rand(3,4,generator=g)>0.3).float()
eq("masked_max 2d", R.masked_max(vv,m2), O.masked_max(vv,m2))
print("done")
Edge case tests
python3 p24.py
FAIL bucketize boundary R=tensor([1, 2, 3, 1, 1, 2]) O=tensor([0, 1, 2, 0, 1, 2]) OK top_p p=1.0 OK top_p p=0.0 OK first_nonzero OK masked_mean axis-1 zerorow OK masked_mean axis1 OK quantile 0.1 OK quantile 0.33 OK quantile 0.9 OK quantile 1.0 OK quantile 0.0 FAIL mode tie R=3 O=1 OK ce 2d OK imp None OK gae tensor nv OK gae float nv OK rev_kl_obj OK entropy allzero OK masked_max 2d done [stdout] FAIL bucketize boundary R=tensor([1, 2, 3, 1, 1, 2]) O=tensor([0, 1, 2, 0, 1, 2]) OK top_p p=1.0 OK top_p p=0.0 OK first_nonzero OK masked_mean axis-1 zerorow OK masked_mean axis1 OK quantile 0.1 OK quantile 0.33 OK quantile 0.9 OK quantile 1.0 OK quantile 0.0 FAIL mode tie R=3 O=1 OK ce 2d OK imp None OK gae tensor nv OK gae float nv OK rev_kl_obj OK entropy allzero OK masked_max 2d done
Probe mode tie-breaking
python3 -c "
import torch, reforacle as R
tests=[[3,1,1,3,2],[1,1,2,2],[2,2,1,1],[5,5,0,0,0,5],[4,4,4,1,1,1,7],[0,1,2],[2,1,0]]
for t in tests:
x=torch.tensor(t)
print(t,'-> R.mode=',R.mode_label(x).item(),' torch.mode=',torch.mode(x).values.item())
"[3, 1, 1, 3, 2] -> R.mode= 3 torch.mode= 1 [1, 1, 2, 2] -> R.mode= 2 torch.mode= 1 [2, 2, 1, 1] -> R.mode= 2 torch.mode= 1 [5, 5, 0, 0, 0, 5] -> R.mode= 5 torch.mode= 0 [4, 4, 4, 1, 1, 1, 7] -> R.mode= 4 torch.mode= 1 [0, 1, 2] -> R.mode= 2 torch.mode= 0 [2, 1, 0] -> R.mode= 2 torch.mode= 0 [stdout] [3, 1, 1, 3, 2] -> R.mode= 3 torch.mode= 1 [1, 1, 2, 2] -> R.mode= 2 torch.mode= 1 [2, 2, 1, 1] -> R.mode= 2 torch.mode= 1 [5, 5, 0, 0, 0, 5] -> R.mode= 5 torch.mode= 0 [4, 4, 4, 1, 1, 1, 7] -> R.mode= 4 torch.mode= 1 [0, 1, 2] -> R.mode= 2 torch.mode= 0 [2, 1, 0] -> R.mode= 2 torch.mode= 0
/app/objectives.py
def mode_label(labels):
return torch.mode(labels).valuesdef mode_label(labels):
vals, counts = torch.unique(labels, return_counts=True)
maxc = counts.max()
return vals[counts == maxc].max()/app/objectives.py
return torch.bucketize(rewards, boundaries)
return torch.bucketize(rewards, boundaries, right=True)
Verify mode/bucketize fixes and full validation
python3 -c "
import torch, reforacle as R, objectives as O
tests=[[3,1,1,3,2],[1,1,2,2],[5,5,0,0,0,5],[4,4,4,1,1,1,7],[0,1,2],[9,9,9]]
for t in tests:
x=torch.tensor(t)
print(t, R.mode_label(x).item(), O.mode_label(x).item())
bd=torch.tensor([0.,1.,2.]); rv=torch.tensor([0.,1.,2.,-0.5,0.9999,1.0001,2.5])
print('buck', R.bucketize_reward(rv,bd).tolist(), O.bucketize_reward(rv,bd).tolist())
" && python3 validate.py 2>&1 | tail -4[3, 1, 1, 3, 2] 3 3 [1, 1, 2, 2] 2 2 [5, 5, 0, 0, 0, 5] 5 5 [4, 4, 4, 1, 1, 1, 7] 4 4 [0, 1, 2] 2 2 [9, 9, 9] 9 9 buck [1, 2, 3, 0, 1, 2, 3] [1, 2, 3, 0, 1, 2, 3] ==== SUMMARY ==== tested 49 functions, 0 failing ALL PASS [stdout] [3, 1, 1, 3, 2] 3 3 [1, 1, 2, 2] 2 2 [5, 5, 0, 0, 0, 5] 5 5 [4, 4, 4, 1, 1, 1, 7] 4 4 [0, 1, 2] 2 2 [9, 9, 9] 9 9 buck [1, 2, 3, 0, 1, 2, 3] [1, 2, 3, 0, 1, 2, 3] ==== SUMMARY ==== tested 49 functions, 0 failing ALL PASS
Verify standalone and 25-trial validation
python3 -c "
import ast,sys
src=open('objectives.py').read()
tree=ast.parse(src)
imps=[n.names[0].name for n in ast.walk(tree) if isinstance(n,(ast.Import,))]+[n.module for n in ast.walk(tree) if isinstance(n,ast.ImportFrom)]
print('imports:', imps)
assert 'reforacle' not in str(imps), 'DEPENDS ON REFORACLE!'
print('standalone OK')
"
# run validate with more trials by editing range
sed -i 's/for trial in range(5)/for trial in range(25)/' validate.py
python3 validate.py 2>&1 | tail -4
sed -i 's/for trial in range(25)/for trial in range(5)/' validate.pyimports: ['math', 'torch', 'torch.nn.functional', '__future__'] standalone OK ==== SUMMARY ==== tested 49 functions, 0 failing ALL PASS [stdout] imports: ['math', 'torch', 'torch.nn.functional', '__future__'] standalone OK ==== SUMMARY ==== tested 49 functions, 0 failing ALL PASS
Clean probe files, check budget
rm -f p1.py p2.py p3.py p4.py p5.py p6.py p7.py p8.py p9.py p10.py p11.py p12.py p13.py p14.py p15.py p16.py p17.py p18.py p19.py p20.py p21.py p22.py p23.py p24.py; bash timer.sh
Remaining time (hours:minutes): 1:43 [stdout] Remaining time (hours:minutes): 1:43
No file changes, the agent left the workspace untouched.
=== SyncVals verifier (post-training objectives, sealed oracle) === FUNCTIONS PASSED: 45/49 (partial_score=0.9184) 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_84edb7d3a14c4997. Re-running the agent live requires EVAL_PLATFORM_ENABLE_OAUTH_SMOKE=1 and is non-deterministic.
Trial trial_84edb7d3a14c4997 · verifier authoritative; classifier explanatory.