tests/test.sh exit code (0 → resolved); the classification below is post-hoc and cannot change it./app/objectives.py
1 from __future__ import annotations
2
3 import torch
4
5
6 def masked_mean(values, mask, axis=None):
7 raise NotImplementedError
8
9
10 def masked_sum(values, mask, axis=None):
11 raise NotImplementedError
12
13
14 def logsumexp(x, axis):
15 raise NotImplementedError
16
17
18 def log_softmax(x, axis):
19 raise NotImplementedError
20
21
22 def selective_logprobs(logits, labels, mask):
23 raise NotImplementedError
24
25
26 def token_logprobs(logits, labels):
27 raise NotImplementedError
28
29
30 def sequence_logprob(logits, labels, mask, length_normalize):
31 raise NotImplementedError
32
33
34 def entropy(logits, mask):
35 raise NotImplementedError
36
37
38 def dpo_loss(pc, pr, rc, rr, beta, label_smoothing):
39 raise NotImplementedError
40
41
42 def ipo_loss(pc, pr, rc, rr, beta):
43 raise NotImplementedError
44
45
46 def grpo_advantages(rewards, group_size, scale_by_std):
47 raise NotImplementedError
48
49
50 def gae(rewards, values, next_value, gamma, lam):
51 raise NotImplementedError
52
53
54 def kl_penalty(logp, ref_logp, estimator):
55 raise NotImplementedError
56
57
58 def clipped_pg_loss(logp, old_logp, advantages, mask, clip_low, clip_high):
59 raise NotImplementedError
60
61
62 def value_loss(values, old_values, returns, clip):
63 raise NotImplementedError
64
65
66 def whiten(values, mask, shift_mean):
67 raise NotImplementedError
68
69
70 def dpo_sequence_loss(pc_logits, pr_logits, rc_logits, rr_logits,
71 chosen_labels, rejected_labels, chosen_mask, rejected_mask,
72 beta, label_smoothing):
73 raise NotImplementedError
74
75
76 def grpo_objective(logits, old_logits, ref_logits, labels, completion_mask,
77 rewards, group_size, beta, clip_low, clip_high, scale_by_std,
78 kl_estimator):
79 raise NotImplementedError
80
81
82 def ppo_objective(rewards, values, old_values, logp, old_logp, next_value,
83 gamma, lam, clip_low, clip_high, vf_clip, vf_coef):
84 raise NotImplementedError
85
86
87 def rloo_advantages(rewards, group_size):
88 raise NotImplementedError
89
90
91 def reverse_kl(logp, ref_logp):
92 raise NotImplementedError
93
94
95 def importance_ratio(logp, old_logp, clip):
96 raise NotImplementedError
97
98
99 def discounted_returns(rewards, gamma):
100 raise NotImplementedError
101
102
103 def normalize(x, eps):
104 raise NotImplementedError
105
106
107 def top_p_mask(probs, p):
108 raise NotImplementedError
109
110
111 def smoothed_nll(logits, labels, smoothing):
112 raise NotImplementedError
113
114
115 def bradley_terry_logit(chosen_reward, rejected_reward, beta):
116 raise NotImplementedError
117
118
119 def rloo_objective(logits, old_logits, labels, mask, rewards, group_size, clip_low, clip_high):
120 raise NotImplementedError
121
122
123 def reverse_kl_objective(logits, ref_logits, labels, mask, advantages, beta):
124 raise NotImplementedError
125
126
127 def cross_entropy(logits, labels, ignore_index):
128 raise NotImplementedError
129
130
131 def top_k_mask(logits, k):
132 raise NotImplementedError
133
134
135 def group_mean_baseline(rewards, group_size):
136 raise NotImplementedError
137
138
139 def lambda_returns(rewards, values, next_value, gamma, lam):
140 raise NotImplementedError
141
142
143 def symmetric_kl(logp, ref_logp):
144 raise NotImplementedError
145
146
147 def huber_value_loss(values, returns, delta):
148 raise NotImplementedError
149
150
151 def normalized_entropy(logits, mask):
152 raise NotImplementedError
153
154
155 def clip_fraction(logp, old_logp, clip):
156 raise NotImplementedError
157
158
159 def masked_whiten(values, mask, shift_mean):
160 raise NotImplementedError
161
162
163 def logprob_at_temperature(logits, labels, temperature):
164 raise NotImplementedError
165
166
167 def advantage_mean_std(advantages, mask):
168 raise NotImplementedError
169
170
171 def argmax_tokens(logits):
172 raise NotImplementedError
173
174
175 def mode_label(labels):
176 raise NotImplementedError
177
178
179 def median_reward(rewards):
180 raise NotImplementedError
181
182
183 def masked_max(values, mask):
184 raise NotImplementedError
185
186
187 def quantile_lower(x, q):
188 raise NotImplementedError
189
190
191 def pad_mask_from_lengths(lengths, max_len):
192 raise NotImplementedError
193
194
195 def first_nonzero_index(mask):
196 raise NotImplementedError
197
198
199 def cumulative_max(x):
200 raise NotImplementedError
201
202
203 def bucketize_reward(rewards, boundaries):
204 raise NotImplementedError
205
1 from __future__ import annotations
2
3 import torch
4
5 import reforacle
6
7 g = torch.Generator().manual_seed(0)
8 values = torch.randn(8, generator=g)
9 mask = (torch.rand(8, generator=g) > 0.5).float()
10
11 print("inputs:")
12 print(" values =", values.tolist())
13 print(" mask =", mask.tolist())
14 print("reference masked_mean ->", reforacle.masked_mean(values, mask).item())
15 print("reference masked_mean (empty mask) ->", reforacle.masked_mean(torch.randn(4, generator=g), torch.zeros(4)).item())
16 print()
17 print("Probe any reference function this way to discover its exact numerical")
18 print("convention, then implement it in /app/objectives.py to match.")
19
Remaining time (hours:minutes):
1:59
[stdout]
Remaining time (hours:minutes):
1:59/app/probe.py
1 from __future__ import annotations
2
3 import torch
4
5
6 def masked_mean(values, mask, axis=None):
7 raise NotImplementedError
8
9
10 def masked_sum(values, mask, axis=None):
11 raise NotImplementedError
12
13
14 def logsumexp(x, axis):
15 raise NotImplementedError
16
17
18 def log_softmax(x, axis):
19 raise NotImplementedError
20
21
22 def selective_logprobs(logits, labels, mask):
23 raise NotImplementedError
24
25
26 def token_logprobs(logits, labels):
27 raise NotImplementedError
28
29
30 def sequence_logprob(logits, labels, mask, length_normalize):
31 raise NotImplementedError
32
33
34 def entropy(logits, mask):
35 raise NotImplementedError
36
37
38 def dpo_loss(pc, pr, rc, rr, beta, label_smoothing):
39 raise NotImplementedError
40
41
42 def ipo_loss(pc, pr, rc, rr, beta):
43 raise NotImplementedError
44
45
46 def grpo_advantages(rewards, group_size, scale_by_std):
47 raise NotImplementedError
48
49
50 def gae(rewards, values, next_value, gamma, lam):
51 raise NotImplementedError
52
53
54 def kl_penalty(logp, ref_logp, estimator):
55 raise NotImplementedError
56
57
58 def clipped_pg_loss(logp, old_logp, advantages, mask, clip_low, clip_high):
59 raise NotImplementedError
60
61
62 def value_loss(values, old_values, returns, clip):
63 raise NotImplementedError
64
65
66 def whiten(values, mask, shift_mean):
67 raise NotImplementedError
68
69
70 def dpo_sequence_loss(pc_logits, pr_logits, rc_logits, rr_logits,
71 chosen_labels, rejected_labels, chosen_mask, rejected_mask,
72 beta, label_smoothing):
73 raise NotImplementedError
74
75
76 def grpo_objective(logits, old_logits, ref_logits, labels, completion_mask,
77 rewards, group_size, beta, clip_low, clip_high, scale_by_std,
78 kl_estimator):
79 raise NotImplementedError
80
81
82 def ppo_objective(rewards, values, old_values, logp, old_logp, next_value,
83 gamma, lam, clip_low, clip_high, vf_clip, vf_coef):
84 raise NotImplementedError
85
86
87 def rloo_advantages(rewards, group_size):
88 raise NotImplementedError
89
90
91 def reverse_kl(logp, ref_logp):
92 raise NotImplementedError
93
94
95 def importance_ratio(logp, old_logp, clip):
96 raise NotImplementedError
97
98
99 def discounted_returns(rewards, gamma):
100 raise NotImplementedError
101
102
103 def normalize(x, eps):
104 raise NotImplementedError
105
106
107 def top_p_mask(probs, p):
108 raise NotImplementedError
109
110
111 def smoothed_nll(logits, labels, smoothing):
112 raise NotImplementedError
113
114
115 def bradley_terry_logit(chosen_reward, rejected_reward, beta):
116 raise NotImplementedError
117
118
119 def rloo_objective(logits, old_logits, labels, mask, rewards, group_size, clip_low, clip_high):
120 raise NotImplementedError
121
122
123 def reverse_kl_objective(logits, ref_logits, labels, mask, advantages, beta):
124 raise NotImplementedError
125
126
127 def cross_entropy(logits, labels, ignore_index):
128 raise NotImplementedError
129
130
131 def top_k_mask(logits, k):
132 raise NotImplementedError
133
134
135 def group_mean_baseline(rewards, group_size):
136 raise NotImplementedError
137
138
139 def lambda_returns(rewards, values, next_value, gamma, lam):
140 raise NotImplementedError
141
142
143 def symmetric_kl(logp, ref_logp):
144 raise NotImplementedError
145
146
147 def huber_value_loss(values, returns, delta):
148 raise NotImplementedError
149
150
151 def normalized_entropy(logits, mask):
152 raise NotImplementedError
153
154
155 def clip_fraction(logp, old_logp, clip):
156 raise NotImplementedError
157
158
159 def masked_whiten(values, mask, shift_mean):
160 raise NotImplementedError
161
162
163 def logprob_at_temperature(logits, labels, temperature):
164 raise NotImplementedError
165
166
167 def advantage_mean_std(advantages, mask):
168 raise NotImplementedError
169
170
171 def argmax_tokens(logits):
172 raise NotImplementedError
173
174
175 def mode_label(labels):
176 raise NotImplementedError
177
178
179 def median_reward(rewards):
180 raise NotImplementedError
181
182
183 def masked_max(values, mask):
184 raise NotImplementedError
185
186
187 def quantile_lower(x, q):
188 raise NotImplementedError
189
190
191 def pad_mask_from_lengths(lengths, max_len):
192 raise NotImplementedError
193
194
195 def first_nonzero_index(mask):
196 raise NotImplementedError
197
198
199 def cumulative_max(x):
200 raise NotImplementedError
201
202
203 def bucketize_reward(rewards, boundaries):
204 raise NotImplementedError
205
1 from __future__ import annotations
2
3 import torch
4
5 import reforacle
6
7 g = torch.Generator().manual_seed(0)
8 values = torch.randn(8, generator=g)
9 mask = (torch.rand(8, generator=g) > 0.5).float()
10
11 print("inputs:")
12 print(" values =", values.tolist())
13 print(" mask =", mask.tolist())
14 print("reference masked_mean ->", reforacle.masked_mean(values, mask).item())
15 print("reference masked_mean (empty mask) ->", reforacle.masked_mean(torch.randn(4, generator=g), torch.zeros(4)).item())
16 print()
17 print("Probe any reference function this way to discover its exact numerical")
18 print("convention, then implement it in /app/objectives.py to match.")
19
Remaining time (hours:minutes):
1:59
[stdout]
Remaining time (hours:minutes):
1:59Check remaining budget
bash /app/timer.sh
1 from __future__ import annotations
2
3 import torch
4
5
6 def masked_mean(values, mask, axis=None):
7 raise NotImplementedError
8
9
10 def masked_sum(values, mask, axis=None):
11 raise NotImplementedError
12
13
14 def logsumexp(x, axis):
15 raise NotImplementedError
16
17
18 def log_softmax(x, axis):
19 raise NotImplementedError
20
21
22 def selective_logprobs(logits, labels, mask):
23 raise NotImplementedError
24
25
26 def token_logprobs(logits, labels):
27 raise NotImplementedError
28
29
30 def sequence_logprob(logits, labels, mask, length_normalize):
31 raise NotImplementedError
32
33
34 def entropy(logits, mask):
35 raise NotImplementedError
36
37
38 def dpo_loss(pc, pr, rc, rr, beta, label_smoothing):
39 raise NotImplementedError
40
41
42 def ipo_loss(pc, pr, rc, rr, beta):
43 raise NotImplementedError
44
45
46 def grpo_advantages(rewards, group_size, scale_by_std):
47 raise NotImplementedError
48
49
50 def gae(rewards, values, next_value, gamma, lam):
51 raise NotImplementedError
52
53
54 def kl_penalty(logp, ref_logp, estimator):
55 raise NotImplementedError
56
57
58 def clipped_pg_loss(logp, old_logp, advantages, mask, clip_low, clip_high):
59 raise NotImplementedError
60
61
62 def value_loss(values, old_values, returns, clip):
63 raise NotImplementedError
64
65
66 def whiten(values, mask, shift_mean):
67 raise NotImplementedError
68
69
70 def dpo_sequence_loss(pc_logits, pr_logits, rc_logits, rr_logits,
71 chosen_labels, rejected_labels, chosen_mask, rejected_mask,
72 beta, label_smoothing):
73 raise NotImplementedError
74
75
76 def grpo_objective(logits, old_logits, ref_logits, labels, completion_mask,
77 rewards, group_size, beta, clip_low, clip_high, scale_by_std,
78 kl_estimator):
79 raise NotImplementedError
80
81
82 def ppo_objective(rewards, values, old_values, logp, old_logp, next_value,
83 gamma, lam, clip_low, clip_high, vf_clip, vf_coef):
84 raise NotImplementedError
85
86
87 def rloo_advantages(rewards, group_size):
88 raise NotImplementedError
89
90
91 def reverse_kl(logp, ref_logp):
92 raise NotImplementedError
93
94
95 def importance_ratio(logp, old_logp, clip):
96 raise NotImplementedError
97
98
99 def discounted_returns(rewards, gamma):
100 raise NotImplementedError
101
102
103 def normalize(x, eps):
104 raise NotImplementedError
105
106
107 def top_p_mask(probs, p):
108 raise NotImplementedError
109
110
111 def smoothed_nll(logits, labels, smoothing):
112 raise NotImplementedError
113
114
115 def bradley_terry_logit(chosen_reward, rejected_reward, beta):
116 raise NotImplementedError
117
118
119 def rloo_objective(logits, old_logits, labels, mask, rewards, group_size, clip_low, clip_high):
120 raise NotImplementedError
121
122
123 def reverse_kl_objective(logits, ref_logits, labels, mask, advantages, beta):
124 raise NotImplementedError
125
126
127 def cross_entropy(logits, labels, ignore_index):
128 raise NotImplementedError
129
130
131 def top_k_mask(logits, k):
132 raise NotImplementedError
133
134
135 def group_mean_baseline(rewards, group_size):
136 raise NotImplementedError
137
138
139 def lambda_returns(rewards, values, next_value, gamma, lam):
140 raise NotImplementedError
141
142
143 def symmetric_kl(logp, ref_logp):
144 raise NotImplementedError
145
146
147 def huber_value_loss(values, returns, delta):
148 raise NotImplementedError
149
150
151 def normalized_entropy(logits, mask):
152 raise NotImplementedError
153
154
155 def clip_fraction(logp, old_logp, clip):
156 raise NotImplementedError
157
158
159 def masked_whiten(values, mask, shift_mean):
160 raise NotImplementedError
161
162
163 def logprob_at_temperature(logits, labels, temperature):
164 raise NotImplementedError
165
166
167 def advantage_mean_std(advantages, mask):
168 raise NotImplementedError
169
170
171 def argmax_tokens(logits):
172 raise NotImplementedError
173
174
175 def mode_label(labels):
176 raise NotImplementedError
177
178
179 def median_reward(rewards):
180 raise NotImplementedError
181
182
183 def masked_max(values, mask):
184 raise NotImplementedError
185
186
187 def quantile_lower(x, q):
188 raise NotImplementedError
189
190
191 def pad_mask_from_lengths(lengths, max_len):
192 raise NotImplementedError
193
194
195 def first_nonzero_index(mask):
196 raise NotImplementedError
197
198
199 def cumulative_max(x):
200 raise NotImplementedError
201
202
203 def bucketize_reward(rewards, boundaries):
204 raise NotImplementedError
205
1 from __future__ import annotations
2
3 import torch
4
5 import reforacle
6
7 g = torch.Generator().manual_seed(0)
8 values = torch.randn(8, generator=g)
9 mask = (torch.rand(8, generator=g) > 0.5).float()
10
11 print("inputs:")
12 print(" values =", values.tolist())
13 print(" mask =", mask.tolist())
14 print("reference masked_mean ->", reforacle.masked_mean(values, mask).item())
15 print("reference masked_mean (empty mask) ->", reforacle.masked_mean(torch.randn(4, generator=g), torch.zeros(4)).item())
16 print()
17 print("Probe any reference function this way to discover its exact numerical")
18 print("convention, then implement it in /app/objectives.py to match.")
19
Remaining time (hours:minutes):
1:59
[stdout]
Remaining time (hours:minutes):
1:59List reforacle functions
python3 -c "import reforacle; print([x for x in dir(reforacle) if not x.startswith('_')])"['F', 'advantage_mean_std', 'argmax_tokens', 'bradley_terry_logit', 'bucketize_reward', 'clip_fraction', 'clipped_pg_loss', 'cross_entropy', 'cumulative_max', 'discounted_returns', 'dpo_loss', 'dpo_sequence_loss', 'entropy', 'first_nonzero_index', 'gae', 'group_mean_baseline', 'grpo_advantages', 'grpo_objective', 'huber_value_loss', 'importance_ratio', 'ipo_loss', 'kl_penalty', 'lambda_returns', 'log_softmax', 'logprob_at_temperature', 'logsumexp', 'masked_max', 'masked_mean', 'masked_sum', 'masked_whiten', 'median_reward', 'mode_label', 'normalize', 'normalized_entropy', 'pad_mask_from_lengths', 'ppo_objective', 'quantile_lower', 'reverse_kl', 'reverse_kl_objective', 'rloo_advantages', 'rloo_objective', 'selective_logprobs', 'sequence_logprob', 'smoothed_nll', 'symmetric_kl', 'token_logprobs', 'top_k_mask', 'top_p_mask', 'torch', 'value_loss', 'whiten'] [stdout] ['F', 'advantage_mean_std', 'argmax_tokens', 'bradley_terry_logit', 'bucketize_reward', 'clip_fraction', 'clipped_pg_loss', 'cross_entropy', 'cumulative_max', 'discounted_returns', 'dpo_loss', 'dpo_sequence_loss', 'entropy', 'first_nonzero_index', 'gae', 'group_mean_baseline', 'grpo_advantages', 'grpo_objective', 'huber_value_loss', 'importance_ratio', 'ipo_loss', 'kl_penalty', 'lambda_returns', 'log_softmax', 'logprob_at_temperature', 'logsumexp', 'masked_max', 'masked_mean', 'masked_sum', 'masked_whiten', 'median_reward', 'mode_label', 'normalize', 'normalized_entropy', 'pad_mask_from_lengths', 'ppo_objective', 'quantile_lower', 'reverse_kl', 'reverse_kl_objective', 'rloo_advantages', 'rloo_objective', 'selective_logprobs', 'sequence_logprob', 'smoothed_nll', 'symmetric_kl', 'token_logprobs', 'top_k_mask', 'top_p_mask', 'torch', 'value_loss', 'whiten']
Probe masked_mean/sum
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(1)
v=torch.randn(8,generator=g); m=(torch.rand(8,generator=g)>0.5).float()
print('masked_mean', reforacle.masked_mean(v,m).item())
print('manual', (v*m).sum().item()/m.sum().item())
print('masked_sum', reforacle.masked_sum(v,m).item(), (v*m).sum().item())
# 2d with axis
v2=torch.randn(3,4,generator=g); m2=(torch.rand(3,4,generator=g)>0.4).float()
print('mm axis0', reforacle.masked_mean(v2,m2,axis=0))
print('manual', (v2*m2).sum(0)/m2.sum(0))
print('mm axis1', reforacle.masked_mean(v2,m2,axis=1))
print('ms axis1', reforacle.masked_sum(v2,m2,axis=1))
# empty mask handling
print('empty', reforacle.masked_mean(torch.randn(4),torch.zeros(4)).item())
"
masked_mean 0.44811558723449707 manual 0.44811558723449707 masked_sum 1.3443467617034912 1.3443467617034912 mm axis0 tensor([ 1.5829, -0.9656, 0.0130, -0.7001]) manual tensor([ 1.5829, -0.9656, 0.0130, -0.7001]) mm axis1 tensor([-0.9656, -0.7057, 1.3029]) ms axis1 tensor([-0.9656, -2.1171, 3.9087]) empty 0.0 [stdout] masked_mean 0.44811558723449707 manual 0.44811558723449707 masked_sum 1.3443467617034912 1.3443467617034912 mm axis0 tensor([ 1.5829, -0.9656, 0.0130, -0.7001]) manual tensor([ 1.5829, -0.9656, 0.0130, -0.7001]) mm axis1 tensor([-0.9656, -0.7057, 1.3029]) ms axis1 tensor([-0.9656, -2.1171, 3.9087]) empty 0.0
Probe logsumexp, softmax, token_logprobs
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(2)
x=torch.randn(3,5,generator=g)
print('logsumexp ax1', reforacle.logsumexp(x,1), torch.logsumexp(x,1))
print('log_softmax ax1 match', torch.allclose(reforacle.log_softmax(x,1), torch.log_softmax(x,1)))
# token_logprobs
logits=torch.randn(2,4,6,generator=g); labels=torch.randint(0,6,(2,4),generator=g)
tl=reforacle.token_logprobs(logits,labels)
print('token_logprobs shape',tl.shape)
lp=torch.log_softmax(logits,-1)
man=lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
print('token match', torch.allclose(tl,man))
mask=(torch.rand(2,4,generator=g)>0.3).float()
print('selective', reforacle.selective_logprobs(logits,labels,mask))
print('sel manual', man*mask)
"
logsumexp ax1 tensor([1.8161, 1.6195, 1.9888]) tensor([1.8161, 1.6195, 1.9888])
log_softmax ax1 match True
token_logprobs shape torch.Size([2, 4])
token match True
selective tensor([ -1.4438, -11.4965])
sel manual tensor([[-0.0000, -0.0000, -1.4438, -0.0000],
[-2.2239, -3.2668, -3.0959, -2.9099]])
[stdout]
logsumexp ax1 tensor([1.8161, 1.6195, 1.9888]) tensor([1.8161, 1.6195, 1.9888])
log_softmax ax1 match True
token_logprobs shape torch.Size([2, 4])
token match True
selective tensor([ -1.4438, -11.4965])
sel manual tensor([[-0.0000, -0.0000, -1.4438, -0.0000],
[-2.2239, -3.2668, -3.0959, -2.9099]])Probe sequence_logprob and entropy
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(3)
logits=torch.randn(2,4,6,generator=g); labels=torch.randint(0,6,(2,4),generator=g)
mask=(torch.rand(2,4,generator=g)>0.3).float()
lp=torch.log_softmax(logits,-1); tok=lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
print('sel', reforacle.selective_logprobs(logits,labels,mask), (tok*mask).sum(-1))
print('seq norm F', reforacle.sequence_logprob(logits,labels,mask,False))
print('seq norm T', reforacle.sequence_logprob(logits,labels,mask,True))
print('sum', (tok*mask).sum(-1), 'div', (tok*mask).sum(-1)/mask.sum(-1))
# entropy
e=reforacle.entropy(logits,mask)
print('entropy',e)
p=torch.softmax(logits,-1); ent=-(p*lp).sum(-1)
print('ent per tok', ent)
print('masked mean ent', (ent*mask).sum(-1)/mask.sum(-1))
print('masked sum ent',(ent*mask).sum(-1))
"
sel tensor([-5.7723, -1.1747]) tensor([-5.7723, -1.1747])
seq norm F tensor([-5.7723, -1.1747])
seq norm T tensor([-1.9241, -0.5874])
sum tensor([-5.7723, -1.1747]) div tensor([-1.9241, -0.5874])
entropy tensor(1.4465)
ent per tok tensor([[1.6965, 1.3433, 1.7497, 1.5413],
[1.4232, 1.0196, 1.7264, 1.6772]])
masked mean ent tensor([1.5965, 1.2214])
masked sum ent tensor([4.7894, 2.4428])
[stdout]
sel tensor([-5.7723, -1.1747]) tensor([-5.7723, -1.1747])
seq norm F tensor([-5.7723, -1.1747])
seq norm T tensor([-1.9241, -0.5874])
sum tensor([-5.7723, -1.1747]) div tensor([-1.9241, -0.5874])
entropy tensor(1.4465)
ent per tok tensor([[1.6965, 1.3433, 1.7497, 1.5413],
[1.4232, 1.0196, 1.7264, 1.6772]])
masked mean ent tensor([1.5965, 1.2214])
masked sum ent tensor([4.7894, 2.4428])Entropy scalar check
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(3)
logits=torch.randn(2,4,6,generator=g); labels=torch.randint(0,6,(2,4),generator=g)
mask=(torch.rand(2,4,generator=g)>0.3).float()
lp=torch.log_softmax(logits,-1)
p=torch.softmax(logits,-1); ent=-(p*lp).sum(-1)
print('global masked mean', (ent*mask).sum()/mask.sum())
print('ref', reforacle.entropy(logits,mask))
# no mask?
print('entropy no mask via ones', reforacle.entropy(logits, torch.ones(2,4)))
print('mean ent all', ent.mean())
"
global masked mean tensor(1.4465) ref tensor(1.4465) entropy no mask via ones tensor(1.5222) mean ent all tensor(1.5222) [stdout] global masked mean tensor(1.4465) ref tensor(1.4465) entropy no mask via ones tensor(1.5222) mean ent all tensor(1.5222)
Probe dpo and ipo
python3 -c "
import torch, reforacle
import torch.nn.functional as F
g=torch.Generator().manual_seed(4)
pc=torch.randn(5,generator=g); pr=torch.randn(5,generator=g); rc=torch.randn(5,generator=g); rr=torch.randn(5,generator=g)
beta=0.1; ls=0.0
d=reforacle.dpo_loss(pc,pr,rc,rr,beta,ls)
print('dpo',d)
logits=beta*((pc-rc)-(pr-rr))
print('mean -logsig', (-F.logsigmoid(logits)).mean())
# with label smoothing
d2=reforacle.dpo_loss(pc,pr,rc,rr,beta,0.2)
print('dpo ls0.2', d2)
ls=0.2
man=(-F.logsigmoid(logits)*(1-ls)-F.logsigmoid(-logits)*ls).mean()
print('man ls', man)
# ipo
i=reforacle.ipo_loss(pc,pr,rc,rr,beta)
print('ipo', i)
print('ipo man', ((logits-1/(2*beta))**2).mean())
"
dpo tensor(0.7024) mean -logsig tensor(0.7024) dpo ls0.2 tensor(0.7003) man ls tensor(0.7003) ipo tensor(29.2614) ipo man tensor(25.1378) [stdout] dpo tensor(0.7024) mean -logsig tensor(0.7024) dpo ls0.2 tensor(0.7003) man ls tensor(0.7003) ipo tensor(29.2614) ipo man tensor(25.1378)
ipo variants
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(4)
pc=torch.randn(5,generator=g); pr=torch.randn(5,generator=g); rc=torch.randn(5,generator=g); rr=torch.randn(5,generator=g)
beta=0.1
h=(pc-rc)-(pr-rr)
print('ipo', reforacle.ipo_loss(pc,pr,rc,rr,beta))
print('man1', ((h-1/(2*beta))**2).mean())
print('man2', ((beta*h-1/(2*beta))**2).mean())
# try beta out front differently: (h - 1/(2beta))^2 gave?
print('man3', (beta*(h-1/(2*beta))**2).mean())
"
ipo tensor(29.2614) man1 tensor(29.2614) man2 tensor(25.1378) man3 tensor(2.9261) [stdout] ipo tensor(29.2614) man1 tensor(29.2614) man2 tensor(25.1378) man3 tensor(2.9261)
grpo_advantages
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(5)
rewards=torch.randn(6,generator=g)
print('grpo gs3 std=T', reforacle.grpo_advantages(rewards,3,True))
print('grpo gs3 std=F', reforacle.grpo_advantages(rewards,3,False))
r=rewards.view(2,3)
mean=r.mean(1,keepdim=True); std=r.std(1,keepdim=True)
print('man std T', ((r-mean)/(std+1e-6)).view(-1))
print('man std F', (r-mean).view(-1))
# unbiased std?
std2=r.std(1,keepdim=True,unbiased=False)
print('man std T unbiased=F', ((r-mean)/(std2+1e-6)).view(-1))
"
grpo gs3 std=T tensor([ 1.3036, -1.1265, -0.1771, -0.1556, -1.1395, 1.2951]) grpo gs3 std=F tensor([ 0.0628, -0.0543, -0.0085, -0.1368, -1.0017, 1.1385]) man std T tensor([ 1.0644, -0.9198, -0.1446, -0.1270, -0.9304, 1.0574]) man std F tensor([ 0.0628, -0.0543, -0.0085, -0.1368, -1.0017, 1.1385]) man std T unbiased=F tensor([ 1.3036, -1.1265, -0.1771, -0.1556, -1.1395, 1.2951]) [stdout] grpo gs3 std=T tensor([ 1.3036, -1.1265, -0.1771, -0.1556, -1.1395, 1.2951]) grpo gs3 std=F tensor([ 0.0628, -0.0543, -0.0085, -0.1368, -1.0017, 1.1385]) man std T tensor([ 1.0644, -0.9198, -0.1446, -0.1270, -0.9304, 1.0574]) man std F tensor([ 0.0628, -0.0543, -0.0085, -0.1368, -1.0017, 1.1385]) man std T unbiased=F tensor([ 1.3036, -1.1265, -0.1771, -0.1556, -1.1395, 1.2951])
gae
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(6)
rewards=torch.randn(2,4,generator=g); values=torch.randn(2,4,generator=g); next_value=torch.randn(2,generator=g)
gamma=0.99; lam=0.95
adv=reforacle.gae(rewards,values,next_value,gamma,lam)
print('gae',adv)
# manual gae
T=4
def man_gae(rewards,values,next_value):
adv=torch.zeros_like(rewards)
lastgae=torch.zeros(rewards.shape[0])
for t in reversed(range(T)):
nv = values[:,t+1] if t+1<T else next_value
delta=rewards[:,t]+gamma*nv-values[:,t]
lastgae=delta+gamma*lam*lastgae
adv[:,t]=lastgae
return adv
print('man',man_gae(rewards,values,next_value))
"
Exit code 1 Traceback (most recent call last): File "<string>", line 6, in <module> File "reforacle.py", line 82, in reforacle.gae RuntimeError: The size of tensor a (4) must match the size of tensor b (2) at non-singleton dimension 0 [error] tool reported failure
gae 1d, lambda_returns, discounted_returns
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(6)
rewards=torch.randn(4,generator=g); values=torch.randn(4,generator=g); next_value=torch.tensor(0.5)
gamma=0.99; lam=0.95
adv=reforacle.gae(rewards,values,next_value,gamma,lam)
print('gae',adv)
T=4
def man_gae(rewards,values,next_value):
adv=torch.zeros_like(rewards)
lastgae=0.0
for t in reversed(range(T)):
nv = values[t+1] if t+1<T else next_value
delta=rewards[t]+gamma*nv-values[t]
lastgae=delta+gamma*lam*lastgae
adv[t]=lastgae
return adv
print('man',man_gae(rewards,values,next_value))
print('lambda_returns', reforacle.lambda_returns(rewards,values,next_value,gamma,lam))
print('gae+values', adv+values)
print('discounted_returns', reforacle.discounted_returns(rewards,gamma))
dr=torch.zeros(4); acc=0
for t in reversed(range(4)):
acc=rewards[t]+gamma*acc; dr[t]=acc
print('man dr', dr)
"
gae tensor([-0.7665, -0.5737, -0.8969, -0.8438]) man tensor([-0.7665, -0.5737, -0.8969, -0.8438]) lambda_returns tensor([-2.2381, -0.3960, 0.5588, -0.2035]) gae+values tensor([-2.2381, -0.3960, 0.5588, -0.2035]) discounted_returns tensor([-2.8317, -0.9669, 0.0270, -0.6985]) man dr tensor([-2.8317, -0.9669, 0.0270, -0.6985]) [stdout] gae tensor([-0.7665, -0.5737, -0.8969, -0.8438]) man tensor([-0.7665, -0.5737, -0.8969, -0.8438]) lambda_returns tensor([-2.2381, -0.3960, 0.5588, -0.2035]) gae+values tensor([-2.2381, -0.3960, 0.5588, -0.2035]) discounted_returns tensor([-2.8317, -0.9669, 0.0270, -0.6985]) man dr tensor([-2.8317, -0.9669, 0.0270, -0.6985])
kl estimators
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(7)
logp=torch.randn(6,generator=g); ref=torch.randn(6,generator=g)
for est in ['k1','k2','k3','kl','mse','abs']:
try:
print(est, reforacle.kl_penalty(logp,ref,est))
except Exception as e:
print(est,'ERR',str(e)[:40])
d=logp-ref
print('k1=logp-ref', d)
print('k2=.5 d^2', 0.5*d**2)
print('k3', (ref-logp).exp()-1-(ref-logp))
print('reverse_kl', reforacle.reverse_kl(logp,ref))
print('symmetric_kl', reforacle.symmetric_kl(logp,ref))
"
k1 tensor([ 0.2088, -0.4462, 0.8086, 0.5679, 1.3731, -1.0276]) k2 tensor([0.0218, 0.0996, 0.3270, 0.1612, 0.9427, 0.5280]) k3 tensor([0.0204, 0.1162, 0.2541, 0.1346, 0.6264, 0.7668]) kl ERR kl mse ERR mse abs ERR abs k1=logp-ref tensor([ 0.2088, -0.4462, 0.8086, 0.5679, 1.3731, -1.0276]) k2=.5 d^2 tensor([0.0218, 0.0996, 0.3270, 0.1612, 0.9427, 0.5280]) k3 tensor([0.0204, 0.1162, 0.2541, 0.1346, 0.6264, 0.7668]) reverse_kl tensor([0.0204, 0.1162, 0.2541, 0.1346, 0.6264, 0.7668]) symmetric_kl tensor([0.0219, 0.1012, 0.3452, 0.1656, 1.1005, 0.5762]) [stdout] k1 tensor([ 0.2088, -0.4462, 0.8086, 0.5679, 1.3731, -1.0276]) k2 tensor([0.0218, 0.0996, 0.3270, 0.1612, 0.9427, 0.5280]) k3 tensor([0.0204, 0.1162, 0.2541, 0.1346, 0.6264, 0.7668]) kl ERR kl mse ERR mse abs ERR abs k1=logp-ref tensor([ 0.2088, -0.4462, 0.8086, 0.5679, 1.3731, -1.0276]) k2=.5 d^2 tensor([0.0218, 0.0996, 0.3270, 0.1612, 0.9427, 0.5280]) k3 tensor([0.0204, 0.1162, 0.2541, 0.1346, 0.6264, 0.7668]) reverse_kl tensor([0.0204, 0.1162, 0.2541, 0.1346, 0.6264, 0.7668]) symmetric_kl tensor([0.0219, 0.1012, 0.3452, 0.1656, 1.1005, 0.5762])
symmetric_kl
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(7)
logp=torch.randn(6,generator=g); ref=torch.randn(6,generator=g)
sk=reforacle.symmetric_kl(logp,ref)
print('sk',sk)
d=logp-ref
rkl=(ref-logp).exp()-1-(ref-logp) # reverse
fkl=(logp-ref).exp()-1-(logp-ref) # forward
print('0.5(f+r)',0.5*(fkl+rkl))
print('f+r',fkl+rkl)
print('rkl+fkl each', rkl, fkl)
# maybe symmetric = 0.5*(exp(d)-exp(-d))*d ...
# try (exp(d)-1)*d?
print('(exp(d)-1)*d/... ', )
"
sk tensor([0.0219, 0.1012, 0.3452, 0.1656, 1.1005, 0.5762]) 0.5(f+r) tensor([0.0219, 0.1012, 0.3452, 0.1656, 1.1005, 0.5762]) f+r tensor([0.0438, 0.2025, 0.6903, 0.3312, 2.2009, 1.1523]) rkl+fkl each tensor([0.0204, 0.1162, 0.2541, 0.1346, 0.6264, 0.7668]) tensor([0.0234, 0.0863, 0.4362, 0.1966, 1.5745, 0.3855]) (exp(d)-1)*d/... [stdout] sk tensor([0.0219, 0.1012, 0.3452, 0.1656, 1.1005, 0.5762]) 0.5(f+r) tensor([0.0219, 0.1012, 0.3452, 0.1656, 1.1005, 0.5762]) f+r tensor([0.0438, 0.2025, 0.6903, 0.3312, 2.2009, 1.1523]) rkl+fkl each tensor([0.0204, 0.1162, 0.2541, 0.1346, 0.6264, 0.7668]) tensor([0.0234, 0.0863, 0.4362, 0.1966, 1.5745, 0.3855]) (exp(d)-1)*d/...
importance_ratio, clip_fraction, normalize
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(8)
logp=torch.randn(6,generator=g); old=torch.randn(6,generator=g)
print('imp clip None', reforacle.importance_ratio(logp,old,None))
print('man', (logp-old).exp())
print('imp clip 0.2', reforacle.importance_ratio(logp,old,0.2))
print('man clip', (logp-old).exp().clamp(1-0.2,1+0.2))
print('clip_fraction', reforacle.clip_fraction(logp,old,0.2))
r=(logp-old).exp()
print('frac man', ((r>1.2)|(r<0.8)).float().mean())
print('normalize', reforacle.normalize(torch.randn(5,generator=g),1e-8))
"
imp clip None tensor([1.5138, 0.3108, 7.6442, 0.8076, 2.0245, 0.1235]) man tensor([1.5138, 0.3108, 7.6442, 0.8076, 2.0245, 0.1235]) imp clip 0.2 tensor([1.2000, 0.8000, 1.2000, 0.8076, 1.2000, 0.8000]) man clip tensor([1.2000, 0.8000, 1.2000, 0.8076, 1.2000, 0.8000]) clip_fraction tensor(0.8333) frac man tensor(0.8333) normalize tensor([-1.5946, 1.1741, 0.2035, 0.0250, 0.1920]) [stdout] imp clip None tensor([1.5138, 0.3108, 7.6442, 0.8076, 2.0245, 0.1235]) man tensor([1.5138, 0.3108, 7.6442, 0.8076, 2.0245, 0.1235]) imp clip 0.2 tensor([1.2000, 0.8000, 1.2000, 0.8076, 1.2000, 0.8000]) man clip tensor([1.2000, 0.8000, 1.2000, 0.8076, 1.2000, 0.8000]) clip_fraction tensor(0.8333) frac man tensor(0.8333) normalize tensor([-1.5946, 1.1741, 0.2035, 0.0250, 0.1920])
normalize std bias, whiten
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(20)
x=torch.randn(5,generator=g)
print('normalize', reforacle.normalize(x,1e-8))
for unb in [True,False]:
print('std unb',unb,(x-x.mean())/(x.std(unbiased=unb)+1e-8))
# whiten
g=torch.Generator().manual_seed(21)
v=torch.randn(6,generator=g)
print('whiten shiftT', reforacle.whiten(v,None,True))
print('whiten shiftF', reforacle.whiten(v,None,False))
print('man wT', (v-v.mean())/(v.std(unbiased=False)+1e-8))
print('man wT unb', (v-v.mean())/(v.std(unbiased=True)+1e-8))
"
Exit code 1 normalize tensor([-0.4943, 0.5665, 1.4881, -0.7410, -0.8193]) std unb True tensor([-0.4943, 0.5665, 1.4881, -0.7410, -0.8193]) std unb False tensor([-0.5527, 0.6333, 1.6638, -0.8284, -0.9160]) Traceback (most recent call last): File "<string>", line 11, in <module> 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 failure
whiten with mask
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(21)
v=torch.randn(6,generator=g); m=(torch.rand(6,generator=g)>0.3).float()
print('whiten shiftT', reforacle.whiten(v,m,True))
print('whiten shiftF', reforacle.whiten(v,m,False))
print('masked_whiten shiftT', reforacle.masked_whiten(v,m,True))
# masked mean/var
mean=(v*m).sum()/m.sum()
var=((v-mean)**2*m).sum()/m.sum()
print('shiftT man (v-mean)/sqrt(var+1e-8)', (v-mean)/torch.sqrt(var+1e-8))
print('shiftF man v/sqrt(var+1e-8)+mean?', v/torch.sqrt(var+1e-8))
# whiten all ones mask
mo=torch.ones(6)
print('whiten ones T', reforacle.whiten(v,mo,True))
meanA=v.mean();
print('unbiased var mask?')
"
whiten shiftT tensor([ 1.5110, -0.2302, -1.2899, 0.5505, 0.0091, -7.9123]) whiten shiftF tensor([ 1.1455, -0.5956, -1.6554, 0.1850, -0.3563, -8.2777]) masked_whiten shiftT tensor([ 1.5110, -0.2302, -1.2899, 0.0000, 0.0091, -0.0000]) shiftT man (v-mean)/sqrt(var+1e-8) tensor([ 1.5110, -0.2302, -1.2899, 0.5505, 0.0091, -7.9123]) shiftF man v/sqrt(var+1e-8)+mean? tensor([ 0.3449, -1.3962, -2.4560, -0.6156, -1.1569, -9.0783]) whiten ones T tensor([ 0.8816, 0.3210, -0.0203, 0.5723, 0.3980, -2.1526]) unbiased var mask? [stdout] whiten shiftT tensor([ 1.5110, -0.2302, -1.2899, 0.5505, 0.0091, -7.9123]) whiten shiftF tensor([ 1.1455, -0.5956, -1.6554, 0.1850, -0.3563, -8.2777]) masked_whiten shiftT tensor([ 1.5110, -0.2302, -1.2899, 0.0000, 0.0091, -0.0000]) shiftT man (v-mean)/sqrt(var+1e-8) tensor([ 1.5110, -0.2302, -1.2899, 0.5505, 0.0091, -7.9123]) shiftF man v/sqrt(var+1e-8)+mean? tensor([ 0.3449, -1.3962, -2.4560, -0.6156, -1.1569, -9.0783]) whiten ones T tensor([ 0.8816, 0.3210, -0.0203, 0.5723, 0.3980, -2.1526]) unbiased var mask?
clipped_pg_loss
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(9)
logp=torch.randn(2,4,generator=g); old=torch.randn(2,4,generator=g); adv=torch.randn(2,4,generator=g)
mask=(torch.rand(2,4,generator=g)>0.3).float()
loss=reforacle.clipped_pg_loss(logp,old,adv,mask,0.2,0.2)
print('clipped_pg',loss)
r=(logp-old).exp()
l1=-adv*r; l2=-adv*r.clamp(1-0.2,1+0.2)
pg=torch.maximum(l1,l2)
print('masked mean max', (pg*mask).sum()/mask.sum())
print('mean1', (pg*mask).sum()/mask.sum())
"
clipped_pg tensor(0.5367) masked mean max tensor(0.5367) mean1 tensor(0.5367) [stdout] clipped_pg tensor(0.5367) masked mean max tensor(0.5367) mean1 tensor(0.5367)
value_loss, huber
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(10)
values=torch.randn(6,generator=g); old=torch.randn(6,generator=g); returns=torch.randn(6,generator=g)
print('value_loss clip0.2', reforacle.value_loss(values,old,returns,0.2))
vc=old+(values-old).clamp(-0.2,0.2)
l1=(values-returns)**2; l2=(vc-returns)**2
print('0.5 mean max', 0.5*torch.maximum(l1,l2).mean())
print('mean max no half', torch.maximum(l1,l2).mean())
print('value_loss clip None', reforacle.value_loss(values,old,returns,None))
print('0.5 mean l1', 0.5*l1.mean(), 'l1 mean', l1.mean())
# huber
print('huber d1', reforacle.huber_value_loss(values,returns,1.0))
import torch.nn.functional as F
print('smooth_l1', F.smooth_l1_loss(values,returns,beta=1.0))
print('huber func', F.huber_loss(values,returns,delta=1.0))
"
Exit code 1 value_loss clip0.2 tensor(2.3118) 0.5 mean max tensor(2.3118) mean max no half tensor(4.6236) Traceback (most recent call last): File "<string>", line 10, in <module> File "reforacle.py", line 109, in reforacle.value_loss TypeError: bad operand type for unary -: 'NoneType' [error] tool reported failure
huber
python3 -c "
import torch, reforacle
import torch.nn.functional as F
g=torch.Generator().manual_seed(10)
values=torch.randn(6,generator=g); old=torch.randn(6,generator=g); returns=torch.randn(6,generator=g)
print('huber d1', reforacle.huber_value_loss(values,returns,1.0))
print('huber_loss', F.huber_loss(values,returns,delta=1.0))
print('smooth_l1*? ', F.smooth_l1_loss(values,returns,beta=1.0))
print('huber d0.5', reforacle.huber_value_loss(values,returns,0.5))
print('huber_loss0.5', F.huber_loss(values,returns,delta=0.5))
print('0.5*huber', 0.5*F.huber_loss(values,returns,delta=0.5))
"
huber d1 tensor(1.2284) huber_loss tensor(1.2284) smooth_l1*? tensor(1.2284) huber d0.5 tensor(0.7018) huber_loss0.5 tensor(0.7018) 0.5*huber tensor(0.3509) [stdout] huber d1 tensor(1.2284) huber_loss tensor(1.2284) smooth_l1*? tensor(1.2284) huber d0.5 tensor(0.7018) huber_loss0.5 tensor(0.7018) 0.5*huber tensor(0.3509)
normalized_entropy, smoothed_nll, cross_entropy
python3 -c "
import torch, reforacle
import torch.nn.functional as F
g=torch.Generator().manual_seed(11)
logits=torch.randn(2,4,6,generator=g); labels=torch.randint(0,6,(2,4),generator=g)
mask=(torch.rand(2,4,generator=g)>0.3).float()
print('normalized_entropy', reforacle.normalized_entropy(logits,mask))
lp=torch.log_softmax(logits,-1); p=lp.exp(); ent=-(p*lp).sum(-1)
import math
print('ent/log(V) masked mean', ((ent/math.log(6))*mask).sum()/mask.sum())
print('smoothed_nll s0.1', reforacle.smoothed_nll(logits,labels,0.1))
V=6
nll=-lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
smooth=-lp.mean(-1)
man=(1-0.1)*nll+0.1*smooth
print('man smoothed per', man)
print('man mean', man.mean())
print('cross_entropy', reforacle.cross_entropy(logits,labels,-100))
print('F.ce', F.cross_entropy(logits.view(-1,6),labels.view(-1)))
"
normalized_entropy tensor(0.6777)
ent/log(V) masked mean tensor(0.6777)
smoothed_nll s0.1 tensor(2.6558)
man smoothed per tensor([[2.1796, 3.5302, 2.4966, 2.1949],
[3.6721, 3.1837, 1.7146, 2.2746]])
man mean tensor(2.6558)
cross_entropy tensor(2.6911)
F.ce tensor(2.6911)
[stdout]
normalized_entropy tensor(0.6777)
ent/log(V) masked mean tensor(0.6777)
smoothed_nll s0.1 tensor(2.6558)
man smoothed per tensor([[2.1796, 3.5302, 2.4966, 2.1949],
[3.6721, 3.1837, 1.7146, 2.2746]])
man mean tensor(2.6558)
cross_entropy tensor(2.6911)
F.ce tensor(2.6911)cross_entropy ignore, bradley_terry
python3 -c "
import torch, reforacle
import torch.nn.functional as F
g=torch.Generator().manual_seed(12)
logits=torch.randn(2,4,6,generator=g); labels=torch.randint(0,6,(2,4),generator=g)
labels[0,1]=-100; labels[1,2]=-100
print('cross_entropy', reforacle.cross_entropy(logits,labels,-100))
print('F.ce', F.cross_entropy(logits.view(-1,6),labels.view(-1),ignore_index=-100))
# smoothed with ignore? smoothed_nll has no ignore_index param, skip
# bradley_terry_logit
print('bt', reforacle.bradley_terry_logit(torch.tensor(2.0),torch.tensor(0.5),0.1))
print('man', torch.sigmoid(torch.tensor(0.1*(2.0-0.5))))
print('man logit', 0.1*(2.0-0.5))
"
cross_entropy tensor(1.8114) F.ce tensor(1.8114) bt tensor(0.1500) man tensor(0.5374) man logit 0.15000000000000002 [stdout] cross_entropy tensor(1.8114) F.ce tensor(1.8114) bt tensor(0.1500) man tensor(0.5374) man logit 0.15000000000000002
top_p, top_k, argmax
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(13)
probs=torch.softmax(torch.randn(2,5,generator=g),-1)
print('probs',probs)
print('top_p 0.7', reforacle.top_p_mask(probs,0.7))
print('top_p 0.9', reforacle.top_p_mask(probs,0.9))
logits=torch.randn(2,5,generator=g)
print('logits',logits)
print('top_k 2', reforacle.top_k_mask(logits,2))
print('argmax_tokens', reforacle.argmax_tokens(logits))
"
probs tensor([[0.0727, 0.0259, 0.3333, 0.3186, 0.2494],
[0.0424, 0.0791, 0.2310, 0.5075, 0.1400]])
top_p 0.7 tensor([[0.0000, 0.0000, 0.3698, 0.3535, 0.2767],
[0.0000, 0.0000, 0.3128, 0.6872, 0.0000]])
top_p 0.9 tensor([[0.0000, 0.0000, 0.3698, 0.3535, 0.2767],
[0.0000, 0.0826, 0.2413, 0.5299, 0.1462]])
logits tensor([[ 0.0215, -0.7128, 0.3897, -0.6727, -0.7052],
[-0.5473, 1.2461, -1.4809, -0.3439, 1.9813]])
top_k 2 tensor([[ True, False, True, False, False],
[False, True, False, False, True]])
argmax_tokens tensor([2, 4])
[stdout]
probs tensor([[0.0727, 0.0259, 0.3333, 0.3186, 0.2494],
[0.0424, 0.0791, 0.2310, 0.5075, 0.1400]])
top_p 0.7 tensor([[0.0000, 0.0000, 0.3698, 0.3535, 0.2767],
[0.0000, 0.0000, 0.3128, 0.6872, 0.0000]])
top_p 0.9 tensor([[0.0000, 0.0000, 0.3698, 0.3535, 0.2767],
[0.0000, 0.0826, 0.2413, 0.5299, 0.1462]])
logits tensor([[ 0.0215, -0.7128, 0.3897, -0.6727, -0.7052],
[-0.5473, 1.2461, -1.4809, -0.3439, 1.9813]])
top_k 2 tensor([[ True, False, True, False, False],
[False, True, False, False, True]])
argmax_tokens tensor([2, 4])mode, median, masked_max, quantile
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(14)
labels=torch.tensor([1,3,3,2,3,1])
print('mode_label', reforacle.mode_label(labels))
r=torch.randn(7,generator=g)
print('r',r)
print('median_reward', reforacle.median_reward(r))
print('torch median', r.median())
v=torch.randn(6,generator=g); m=torch.tensor([1.,0,1,0,1,1])
print('masked_max', reforacle.masked_max(v,m))
print('v',v)
print('quantile_lower 0.25', reforacle.quantile_lower(r,0.25))
print('torch.quantile 0.25 lower', torch.quantile(r,0.25,interpolation='lower'))
print('torch.quantile 0.25 linear', torch.quantile(r,0.25))
"
mode_label tensor(3) r tensor([-1.0141, -0.3720, -0.7516, -0.8623, -0.3270, 0.5212, 1.2622]) median_reward tensor(-0.3720) torch median tensor(-0.3720) masked_max tensor(0.5177) v tensor([-1.4680, -0.1037, 0.5177, -1.0845, -2.0901, -0.1508]) quantile_lower 0.25 tensor(-0.8623) torch.quantile 0.25 lower tensor(-0.8623) torch.quantile 0.25 linear tensor(-0.8069) [stdout] mode_label tensor(3) r tensor([-1.0141, -0.3720, -0.7516, -0.8623, -0.3270, 0.5212, 1.2622]) median_reward tensor(-0.3720) torch median tensor(-0.3720) masked_max tensor(0.5177) v tensor([-1.4680, -0.1037, 0.5177, -1.0845, -2.0901, -0.1508]) quantile_lower 0.25 tensor(-0.8623) torch.quantile 0.25 lower tensor(-0.8623) torch.quantile 0.25 linear tensor(-0.8069)
pad_mask, first_nonzero, cummax, bucketize
python3 -c "
import torch, reforacle
lengths=torch.tensor([2,4,1])
print('pad_mask', reforacle.pad_mask_from_lengths(lengths,5))
mask=torch.tensor([[0.,0,1,1,0],[1,1,1,0,0],[0,1,0,0,0]])
print('first_nonzero', reforacle.first_nonzero_index(mask))
x=torch.tensor([1.,3,2,5,4,0])
print('cumulative_max', reforacle.cumulative_max(x))
print('torch cummax', torch.cummax(x,0)[0])
rewards=torch.tensor([0.1,0.5,1.2,-0.3,2.0])
b=torch.tensor([0.0,1.0])
print('bucketize', reforacle.bucketize_reward(rewards,b))
print('torch bucketize right=F', torch.bucketize(rewards,b))
print('torch bucketize right=T', torch.bucketize(rewards,b,right=True))
"
pad_mask tensor([[1., 1., 0., 0., 0.],
[1., 1., 1., 1., 0.],
[1., 0., 0., 0., 0.]])
first_nonzero tensor([2, 0, 1])
cumulative_max tensor([1., 1., 3., 3., 5., 5.])
torch cummax tensor([1., 3., 3., 5., 5., 5.])
bucketize tensor([1, 1, 2, 0, 2])
torch bucketize right=F tensor([1, 1, 2, 0, 2])
torch bucketize right=T tensor([1, 1, 2, 0, 2])
[stdout]
pad_mask tensor([[1., 1., 0., 0., 0.],
[1., 1., 1., 1., 0.],
[1., 0., 0., 0., 0.]])
first_nonzero tensor([2, 0, 1])
cumulative_max tensor([1., 1., 3., 3., 5., 5.])
torch cummax tensor([1., 3., 3., 5., 5., 5.])
bucketize tensor([1, 1, 2, 0, 2])
torch bucketize right=F tensor([1, 1, 2, 0, 2])
torch bucketize right=T tensor([1, 1, 2, 0, 2])cumulative_max more
python3 -c "
import torch, reforacle
for x in [torch.tensor([2.,1,4,3]), torch.tensor([5.,4,3,2,1]), torch.tensor([1.,2,3,4])]:
print(x.tolist(),'->',reforacle.cumulative_max(x).tolist())
# 2d?
x2=torch.tensor([[1.,3,2],[4,1,5]])
try:
print('2d',reforacle.cumulative_max(x2))
except Exception as e: print('2derr',str(e)[:60])
"
[2.0, 1.0, 4.0, 3.0] -> [2.0, 2.0, 2.0, 4.0]
[5.0, 4.0, 3.0, 2.0, 1.0] -> [5.0, 5.0, 5.0, 5.0, 5.0]
[1.0, 2.0, 3.0, 4.0] -> [1.0, 1.0, 2.0, 3.0]
2d tensor([[1., 3., 2.],
[1., 3., 2.]])
[stdout]
[2.0, 1.0, 4.0, 3.0] -> [2.0, 2.0, 2.0, 4.0]
[5.0, 4.0, 3.0, 2.0, 1.0] -> [5.0, 5.0, 5.0, 5.0, 5.0]
[1.0, 2.0, 3.0, 4.0] -> [1.0, 1.0, 2.0, 3.0]
2d tensor([[1., 3., 2.],
[1., 3., 2.]])logprob_at_temp, advantage_mean_std
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(15)
logits=torch.randn(2,4,6,generator=g); labels=torch.randint(0,6,(2,4),generator=g)
print('logprob_at_temp 2.0', reforacle.logprob_at_temperature(logits,labels,2.0))
lp=torch.log_softmax(logits/2.0,-1)
print('man', lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1))
adv=torch.randn(2,4,generator=g); m=(torch.rand(2,4,generator=g)>0.3).float()
print('advantage_mean_std', reforacle.advantage_mean_std(adv,m))
mean=(adv*m).sum()/m.sum(); var=((adv-mean)**2*m).sum()/m.sum()
print('man mean,std biased', mean.item(), var.sqrt().item())
print('man std unbiased', (((adv-mean)**2*m).sum()/(m.sum()-1)).sqrt().item())
"
logprob_at_temp 2.0 tensor([[-1.9129, -1.6168, -2.4023, -1.2461],
[-2.3040, -1.3827, -1.5425, -1.7602]])
man tensor([[-1.9129, -1.6168, -2.4023, -1.2461],
[-2.3040, -1.3827, -1.5425, -1.7602]])
advantage_mean_std tensor([0.2776, 0.6187])
man mean,std biased 0.2775508463382721 0.6187009215354919
man std unbiased 0.6777529120445251
[stdout]
logprob_at_temp 2.0 tensor([[-1.9129, -1.6168, -2.4023, -1.2461],
[-2.3040, -1.3827, -1.5425, -1.7602]])
man tensor([[-1.9129, -1.6168, -2.4023, -1.2461],
[-2.3040, -1.3827, -1.5425, -1.7602]])
advantage_mean_std tensor([0.2776, 0.6187])
man mean,std biased 0.2775508463382721 0.6187009215354919
man std unbiased 0.6777529120445251group_mean_baseline, rloo_advantages
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(16)
rewards=torch.randn(6,generator=g)
print('rewards',rewards)
print('group_mean_baseline gs3', reforacle.group_mean_baseline(rewards,3))
r=rewards.view(2,3)
print('man broadcast', r.mean(1,keepdim=True).expand_as(r).reshape(-1))
print('rloo_advantages gs3', reforacle.rloo_advantages(rewards,3))
# rloo: each - mean of others = (sum-r)/(n-1)
n=3
others=(r.sum(1,keepdim=True)-r)/(n-1)
print('man rloo', (r-others).reshape(-1))
"
rewards tensor([-0.3329, 1.0216, 0.5170, -0.8547, 0.2865, -1.9987])
group_mean_baseline gs3 tensor([-7.3475e-01, 6.1968e-01, 1.1508e-01, 9.3925e-04, 1.1421e+00,
-1.1431e+00])
man broadcast tensor([ 0.4019, 0.4019, 0.4019, -0.8556, -0.8556, -0.8556])
rloo_advantages gs3 tensor([-1.1021e+00, 9.2951e-01, 1.7262e-01, 1.4089e-03, 1.7132e+00,
-1.7146e+00])
man rloo tensor([-1.1021e+00, 9.2951e-01, 1.7262e-01, 1.4089e-03, 1.7132e+00,
-1.7146e+00])
[stdout]
rewards tensor([-0.3329, 1.0216, 0.5170, -0.8547, 0.2865, -1.9987])
group_mean_baseline gs3 tensor([-7.3475e-01, 6.1968e-01, 1.1508e-01, 9.3925e-04, 1.1421e+00,
-1.1431e+00])
man broadcast tensor([ 0.4019, 0.4019, 0.4019, -0.8556, -0.8556, -0.8556])
rloo_advantages gs3 tensor([-1.1021e+00, 9.2951e-01, 1.7262e-01, 1.4089e-03, 1.7132e+00,
-1.7146e+00])
man rloo tensor([-1.1021e+00, 9.2951e-01, 1.7262e-01, 1.4089e-03, 1.7132e+00,
-1.7146e+00])dpo_sequence_loss
python3 -c "
import torch, reforacle
import torch.nn.functional as F
g=torch.Generator().manual_seed(17)
B,T,V=2,4,6
pc=torch.randn(B,T,V,generator=g); pr=torch.randn(B,T,V,generator=g)
rc=torch.randn(B,T,V,generator=g); rr=torch.randn(B,T,V,generator=g)
cl=torch.randint(0,V,(B,T),generator=g); rl=torch.randint(0,V,(B,T),generator=g)
cm=(torch.rand(B,T,generator=g)>0.3).float(); rm=(torch.rand(B,T,generator=g)>0.3).float()
beta=0.1
out=reforacle.dpo_sequence_loss(pc,pr,rc,rr,cl,rl,cm,rm,beta,0.0)
print('out',out)
def seqlp(logits,labels,mask):
lp=torch.log_softmax(logits,-1); tok=lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
return (tok*mask).sum(-1)
pcs=seqlp(pc,cl,cm); prs=seqlp(pr,rl,rm); rcs=seqlp(rc,cl,cm); rrs=seqlp(rr,rl,rm)
logits=beta*((pcs-rcs)-(prs-rrs))
print('man', (-F.logsigmoid(logits)).mean())
"
out tensor(0.8048) man tensor(0.8048) [stdout] out tensor(0.8048) man tensor(0.8048)
grpo_objective
python3 -c "
import torch, reforacle
import torch.nn.functional as F
g=torch.Generator().manual_seed(18)
B,T,V=4,3,5
logits=torch.randn(B,T,V,generator=g); old=torch.randn(B,T,V,generator=g); ref=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g)
cm=(torch.rand(B,T,generator=g)>0.3).float()
rewards=torch.randn(B,generator=g)
gs=2; beta=0.04; cl=0.2; ch=0.2
out=reforacle.grpo_objective(logits,old,ref,labels,cm,rewards,gs,beta,cl,ch,True,'k3')
print('out',out)
# build
def toklp(lg):
lp=torch.log_softmax(lg,-1); return lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
lp=toklp(logits); olp=toklp(old); rlp=toklp(ref)
# advantages
r=rewards.view(-1,gs); mean=r.mean(1,keepdim=True); std=r.std(1,keepdim=True,unbiased=False)
adv=((r-mean)/(std+1e-6)).view(-1) # per sequence
adv=adv.unsqueeze(1) # broadcast over T
ratio=(lp-olp).exp()
l1=-adv*ratio; l2=-adv*ratio.clamp(1-cl,1+ch)
pg=torch.maximum(l1,l2)
kl=(rlp-lp).exp()-1-(rlp-lp)
per=pg+beta*kl
print('masked mean', (per*cm).sum()/cm.sum())
print('pg only', (pg*cm).sum()/cm.sum())
"
out tensor(1.8505) masked mean tensor(1.8505) pg only tensor(1.8205) [stdout] out tensor(1.8505) masked mean tensor(1.8505) pg only tensor(1.8205)
grpo variants
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(19)
B,T,V=4,3,5
logits=torch.randn(B,T,V,generator=g); old=torch.randn(B,T,V,generator=g); ref=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g)
cm=(torch.rand(B,T,generator=g)>0.3).float()
rewards=torch.randn(B,generator=g)
gs=2; beta=0.04; cl=0.2; ch=0.3
def toklp(lg):
lp=torch.log_softmax(lg,-1); return lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
def kl(est,lp,rlp):
d=lp-rlp
if est=='k1': return d
if est=='k2': return 0.5*d**2
if est=='k3': return (rlp-lp).exp()-1-(rlp-lp)
def build(scale,est):
lp=toklp(logits); olp=toklp(old); rlp=toklp(ref)
r=rewards.view(-1,gs); mean=r.mean(1,keepdim=True)
if scale:
std=r.std(1,keepdim=True,unbiased=False); adv=((r-mean)/(std+1e-6)).view(-1)
else: adv=(r-mean).view(-1)
adv=adv.unsqueeze(1)
ratio=(lp-olp).exp(); l1=-adv*ratio; l2=-adv*ratio.clamp(1-cl,1+ch)
pg=torch.maximum(l1,l2); per=pg+beta*kl(est,lp,rlp)
return (per*cm).sum()/cm.sum()
for scale in [True,False]:
for est in ['k1','k2','k3']:
o=reforacle.grpo_objective(logits,old,ref,labels,cm,rewards,gs,beta,cl,ch,scale,est)
print(scale,est,o.item(),build(scale,est).item())
"
True k1 -0.09621106088161469 -0.09621106088161469 True k2 -0.04586809128522873 -0.04586809128522873 True k3 -0.01956360600888729 -0.019563641399145126 False k1 -0.12295626103878021 -0.12295626103878021 False k2 -0.07261323928833008 -0.07261323928833008 False k3 -0.046308789402246475 -0.046308789402246475 [stdout] True k1 -0.09621106088161469 -0.09621106088161469 True k2 -0.04586809128522873 -0.04586809128522873 True k3 -0.01956360600888729 -0.019563641399145126 False k1 -0.12295626103878021 -0.12295626103878021 False k2 -0.07261323928833008 -0.07261323928833008 False k3 -0.046308789402246475 -0.046308789402246475
ppo_objective
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(22)
T=5
rewards=torch.randn(T,generator=g); values=torch.randn(T,generator=g); old_values=torch.randn(T,generator=g)
logp=torch.randn(T,generator=g); old_logp=torch.randn(T,generator=g); next_value=torch.tensor(0.3)
gamma=0.99; lam=0.95; cl=0.2; ch=0.2; vf_clip=0.2; vf_coef=0.5
out=reforacle.ppo_objective(rewards,values,old_values,logp,old_logp,next_value,gamma,lam,cl,ch,vf_clip,vf_coef)
print('out',out)
# gae
def gae(rewards,values,next_value):
adv=torch.zeros_like(rewards); last=0.0
for t in reversed(range(T)):
nv=values[t+1] if t+1<T else next_value
delta=rewards[t]+gamma*nv-values[t]; last=delta+gamma*lam*last; adv[t]=last
return adv
adv=gae(rewards,values,next_value); returns=adv+values
ratio=(logp-old_logp).exp()
l1=-adv*ratio; l2=-adv*ratio.clamp(1-cl,1+ch); pg=torch.maximum(l1,l2).mean()
vc=old_values+(values-old_values).clamp(-vf_clip,vf_clip)
vl=0.5*torch.maximum((values-returns)**2,(vc-returns)**2).mean()
print('pg+vf', pg+vf_coef*vl)
# whiten adv?
advw=(adv-adv.mean())/(adv.std(unbiased=False)+1e-8)
l1=-advw*ratio; l2=-advw*ratio.clamp(1-cl,1+ch); pgw=torch.maximum(l1,l2).mean()
print('with whiten adv', pgw+vf_coef*vl)
advwu=(adv-adv.mean())/(adv.std(unbiased=True)+1e-8)
l1=-advwu*ratio; l2=-advwu*ratio.clamp(1-cl,1+ch); pgwu=torch.maximum(l1,l2).mean()
print('with whiten adv unbiased', pgwu+vf_coef*vl)
"
out tensor(2.0523) pg+vf tensor(5.6713) with whiten adv tensor(2.0523) with whiten adv unbiased tensor(1.9717) [stdout] out tensor(2.0523) pg+vf tensor(5.6713) with whiten adv tensor(2.0523) with whiten adv unbiased tensor(1.9717)
rloo_objective
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(23)
B,T,V=4,3,5
logits=torch.randn(B,T,V,generator=g); old=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g)
mask=(torch.rand(B,T,generator=g)>0.3).float()
rewards=torch.randn(B,generator=g)
gs=2; cl=0.2; ch=0.2
out=reforacle.rloo_objective(logits,old,labels,mask,rewards,gs,cl,ch)
print('out',out)
def toklp(lg):
lp=torch.log_softmax(lg,-1); return lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
lp=toklp(logits); olp=toklp(old)
r=rewards.view(-1,gs); n=gs
adv=(r-(r.sum(1,keepdim=True)-r)/(n-1)).view(-1).unsqueeze(1)
ratio=(lp-olp).exp(); l1=-adv*ratio; l2=-adv*ratio.clamp(1-cl,1+ch)
pg=torch.maximum(l1,l2)
print('masked mean', (pg*mask).sum()/mask.sum())
"
out tensor(0.3387) masked mean tensor(0.3387) [stdout] out tensor(0.3387) masked mean tensor(0.3387)
reverse_kl_objective
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(24)
B,T,V=4,3,5
logits=torch.randn(B,T,V,generator=g); ref=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g)
mask=(torch.rand(B,T,generator=g)>0.3).float()
adv=torch.randn(B,generator=g); beta=0.1
out=reforacle.reverse_kl_objective(logits,ref,labels,mask,adv,beta)
print('out',out)
def toklp(lg):
lp=torch.log_softmax(lg,-1); return lp.gather(-1,labels.unsqueeze(-1)).squeeze(-1)
lp=toklp(logits); rlp=toklp(ref)
kl=(rlp-lp).exp()-1-(rlp-lp)
a=adv.unsqueeze(1)
per=-a*lp+beta*kl
print('v1 -adv*lp+beta*kl', (per*mask).sum()/mask.sum())
per2=-a*lp*mask
print('pg -adv*lp masked', (-a*lp*mask).sum()/mask.sum())
# maybe advantage per token with lp, and kl
# try policy grad -adv*lp + beta*kl
"
out tensor(-0.7979) v1 -adv*lp+beta*kl tensor(-0.7979) pg -adv*lp masked tensor(-0.9014) [stdout] out tensor(-0.7979) v1 -adv*lp+beta*kl tensor(-0.7979) pg -adv*lp masked tensor(-0.9014)
determine eps
python3 -c "
import torch, reforacle
# determine grpo eps: use group with known std
r=torch.tensor([0.0,2.0]) # group of 2, mean1, std biased =1
adv=reforacle.grpo_advantages(r,2,True)
print('adv',adv) # (0-1)/(1+eps), (2-1)/(1+eps)
print('implied denom', 1.0/adv[1].item())
print('eps', 1.0/adv[1].item()-1.0)
# whiten eps
v=torch.tensor([0.0,2.0]); m=torch.ones(2)
w=reforacle.whiten(v,m,True)
print('whiten',w, 'denom', 1.0/w[1].item(), 'eps', (1.0/w[1].item())**2-1)
"
adv tensor([-1.0000, 1.0000]) implied denom 1.000000953675226 eps 9.536752259009518e-07 whiten tensor([-1., 1.]) denom 1.0 eps 0.0 [stdout] adv tensor([-1.0000, 1.0000]) implied denom 1.000000953675226 eps 9.536752259009518e-07 whiten tensor([-1., 1.]) denom 1.0 eps 0.0
edge cases
python3 -c "
import torch, reforacle
# bucketize right? test value exactly on boundary
r=torch.tensor([0.0,1.0,2.0]); b=torch.tensor([0.0,1.0,2.0])
print('bucketize', reforacle.bucketize_reward(r,b))
print('right=F', torch.bucketize(r,b,right=False))
print('right=T', torch.bucketize(r,b,right=True))
# first_nonzero all zero
print('first_nonzero all zero', reforacle.first_nonzero_index(torch.tensor([[0.,0,0]])))
# top_p first token exceeds p
p=torch.tensor([[0.6,0.3,0.1]]); print('top_p 0.5', reforacle.top_p_mask(p,0.5))
# masked_mean multiaxis default
v=torch.randn(3,4); m=torch.ones(3,4)
print('mm None scalar', reforacle.masked_mean(v,m).shape)
"
bucketize tensor([1, 2, 3]) right=F tensor([0, 1, 2]) right=T tensor([1, 2, 3]) first_nonzero all zero tensor([-1]) top_p 0.5 tensor([[1., 0., 0.]]) mm None scalar torch.Size([]) [stdout] bucketize tensor([1, 2, 3]) right=F tensor([0, 1, 2]) right=T tensor([1, 2, 3]) first_nonzero all zero tensor([-1]) top_p 0.5 tensor([[1., 0., 0.]]) mm None scalar torch.Size([])
2d median/mode/quantile
python3 -c "
import torch, reforacle
r=torch.tensor([[1.,5,3],[2,2,9]])
print('median_reward 2d', reforacle.median_reward(r))
print('torch.median global', torch.median(r))
lab=torch.tensor([[1,1,2],[3,3,3]])
print('mode 2d', reforacle.mode_label(lab))
print('quantile 2d', reforacle.quantile_lower(r,0.5))
# even length median
print('median even', reforacle.median_reward(torch.tensor([1.,2,3,4])))
print('masked_max empty mask', reforacle.masked_max(torch.tensor([1.,2,3]),torch.zeros(3)))
"
median_reward 2d tensor([1.5000, 2.5000, 7.0000]) torch.median global tensor(2.) mode 2d tensor([1, 3]) quantile 2d tensor([1., 3., 5.]) median even tensor(2.5000) masked_max empty mask tensor(0.) [stdout] median_reward 2d tensor([1.5000, 2.5000, 7.0000]) torch.median global tensor(2.) mode 2d tensor([1, 3]) quantile 2d tensor([1., 3., 5.]) median even tensor(2.5000) masked_max empty mask tensor(0.)
median hypotheses
python3 -c "
import torch, reforacle
r=torch.tensor([[1.,5,3],[2,2,9]])
print('ref', reforacle.median_reward(r))
print('q0.5 dim0', torch.quantile(r,0.5,dim=0))
print('q0.5 dim1', torch.quantile(r,0.5,dim=1))
print('median dim0', torch.median(r,dim=0)[0])
print('median dim1', torch.median(r,dim=1)[0])
print('mean dim0', r.mean(0))
# 1D checks
for x in [torch.tensor([3.,1,2]), torch.tensor([1.,2,3,4]), torch.tensor([5.,1,3,2,4])]:
print(x.tolist(),'ref',reforacle.median_reward(x).item(),'q',torch.quantile(x,0.5).item(),'med',torch.median(x).item())
"
ref tensor([1.5000, 2.5000, 7.0000]) q0.5 dim0 tensor([1.5000, 3.5000, 6.0000]) q0.5 dim1 tensor([3., 2.]) median dim0 tensor([1., 2., 3.]) median dim1 tensor([3., 2.]) mean dim0 tensor([1.5000, 3.5000, 6.0000]) [3.0, 1.0, 2.0] ref 2.0 q 2.0 med 2.0 [1.0, 2.0, 3.0, 4.0] ref 2.5 q 2.5 med 2.0 [5.0, 1.0, 3.0, 2.0, 4.0] ref 3.0 q 3.0 med 3.0 [stdout] ref tensor([1.5000, 2.5000, 7.0000]) q0.5 dim0 tensor([1.5000, 3.5000, 6.0000]) q0.5 dim1 tensor([3., 2.]) median dim0 tensor([1., 2., 3.]) median dim1 tensor([3., 2.]) mean dim0 tensor([1.5000, 3.5000, 6.0000]) [3.0, 1.0, 2.0] ref 2.0 q 2.0 med 2.0 [1.0, 2.0, 3.0, 4.0] ref 2.5 q 2.5 med 2.0 [5.0, 1.0, 3.0, 2.0, 4.0] ref 3.0 q 3.0 med 3.0
masked_max shape
python3 -c "
import torch, reforacle
v=torch.tensor([[1.,5,3],[2,8,9]]); m=torch.tensor([[1.,0,1],[0,1,0]])
print('masked_max 2d', reforacle.masked_max(v,m))
print('masked_max 1d', reforacle.masked_max(torch.tensor([1.,5,3]),torch.tensor([1.,0,1])))
# masked_max empty per-row?
print('empty', reforacle.masked_max(torch.tensor([1.,2.]),torch.tensor([0.,0.])))
"
masked_max 2d tensor(8.) masked_max 1d tensor(3.) empty tensor(0.) [stdout] masked_max 2d tensor(8.) masked_max 1d tensor(3.) empty tensor(0.)
from __future__ import annotations
import torch
import torch.nn.functional as F
# ---------------------------------------------------------------------------
# Masked reductions
# ---------------------------------------------------------------------------
def masked_mean(values, mask, axis=None):
mask = mask.to(values.dtype)
if axis is None:
num = (values * mask).sum()
den = mask.sum().clamp(min=1.0)
else:
num = (values * mask).sum(axis)
den = mask.sum(axis).clamp(min=1.0)
return num / den
def masked_sum(values, mask, axis=None):
mask = mask.to(values.dtype)
if axis is None:
return (values * mask).sum()
return (values * mask).sum(axis)
def logsumexp(x, axis):
return torch.logsumexp(x, axis)
def log_softmax(x, axis):
return torch.log_softmax(x, axis)
# ---------------------------------------------------------------------------
# Log-probabilities
# ---------------------------------------------------------------------------
def token_logprobs(logits, labels):
lp = torch.log_softmax(logits, dim=-1)
return lp.gather(-1, labels.unsqueeze(-1)).squeeze(-1)
def selective_logprobs(logits, labels, mask):
tok = token_logprobs(logits, labels)
return (tok * mask.to(tok.dtype)).sum(-1)
def sequence_logprob(logits, labels, mask, length_normalize):
tok = token_logprobs(logits, labels)
mask = mask.to(tok.dtype)
s = (tok * mask).sum(-1)
if length_normalize:
return s / mask.sum(-1).clamp(min=1.0)
return s
def logprob_at_temperature(logits, labels, temperature):
lp = torch.log_softmax(logits / temperature, dim=-1)
return lp.gather(-1, labels.unsqueeze(-1)).squeeze(-1)
# ---------------------------------------------------------------------------
# Entropy
# ---------------------------------------------------------------------------
def entropy(logits, mask):
lp = torch.log_softmax(logits, dim=-1)
p = lp.exp()
ent = -(p * lp).sum(-1)
mask = mask.to(ent.dtype)
return (ent * mask).sum() / mask.sum().clamp(min=1.0)
def normalized_entropy(logits, mask):
lp = torch.log_softmax(logits, dim=-1)
p = lp.exp()
ent = -(p * lp).sum(-1)
v = logits.shape[-1]
ent = ent / torch.log(torch.tensor(float(v), dtype=ent.dtype))
mask = mask.to(ent.dtype)
return (ent * mask).sum() / mask.sum().clamp(min=1.0)
# ---------------------------------------------------------------------------
# Preference optimization
# ---------------------------------------------------------------------------
def dpo_loss(pc, pr, rc, rr, beta, label_smoothing):
logits = beta * ((pc - rc) - (pr - rr))
losses = -F.logsigmoid(logits) * (1 - label_smoothing) \
- F.logsigmoid(-logits) * label_smoothing
return losses.mean()
def ipo_loss(pc, pr, rc, rr, beta):
h = (pc - rc) - (pr - rr)
return ((h - 1.0 / (2.0 * beta)) ** 2).mean()
def bradley_terry_logit(chosen_reward, rejected_reward, beta):
return beta * (chosen_reward - rejected_reward)
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)
# ---------------------------------------------------------------------------
# Advantages / baselines
# ---------------------------------------------------------------------------
def grpo_advantages(rewards, group_size, scale_by_std):
r = rewards.view(-1, group_size)
mean = r.mean(1, keepdim=True)
adv = r - mean
if scale_by_std:
std = r.std(1, keepdim=True, unbiased=False)
adv = adv / (std + 1e-6)
return adv.reshape(-1)
def rloo_advantages(rewards, group_size):
r = rewards.view(-1, group_size)
n = group_size
others = (r.sum(1, keepdim=True) - r) / (n - 1)
return (r - others).reshape(-1)
def group_mean_baseline(rewards, group_size):
r = rewards.view(-1, group_size)
mean = r.mean(1, keepdim=True)
return (r - mean).reshape(-1)
# ---------------------------------------------------------------------------
# Returns / advantage estimation
# ---------------------------------------------------------------------------
def gae(rewards, values, next_value, gamma, lam):
T = rewards.shape[0]
adv = torch.zeros_like(rewards)
last = torch.zeros((), dtype=rewards.dtype)
for t in reversed(range(T)):
nv = values[t + 1] if t + 1 < T else next_value
delta = rewards[t] + gamma * nv - values[t]
last = delta + gamma * lam * last
adv[t] = last
return adv
def lambda_returns(rewards, values, next_value, gamma, lam):
return gae(rewards, values, next_value, gamma, lam) + values
def discounted_returns(rewards, gamma):
T = rewards.shape[0]
out = torch.zeros_like(rewards)
acc = torch.zeros((), dtype=rewards.dtype)
for t in reversed(range(T)):
acc = rewards[t] + gamma * acc
out[t] = acc
return out
# ---------------------------------------------------------------------------
# KL divergences
# ---------------------------------------------------------------------------
def kl_penalty(logp, ref_logp, estimator):
d = logp - ref_logp
if estimator == 'k1':
return d
if estimator == 'k2':
return 0.5 * d ** 2
if estimator == 'k3':
return (-d).exp() - 1 - (-d)
raise ValueError(estimator)
def reverse_kl(logp, ref_logp):
d = ref_logp - logp
return d.exp() - 1 - d
def symmetric_kl(logp, ref_logp):
d = logp - ref_logp
fwd = d.exp() - 1 - d
rev = (-d).exp() - 1 - (-d)
return 0.5 * (fwd + rev)
# ---------------------------------------------------------------------------
# Policy-gradient pieces
# ---------------------------------------------------------------------------
def importance_ratio(logp, old_logp, clip):
r = (logp - old_logp).exp()
if clip is not None:
r = r.clamp(1 - clip, 1 + clip)
return r
def clipped_pg_loss(logp, old_logp, advantages, mask, clip_low, clip_high):
ratio = (logp - old_logp).exp()
l1 = -advantages * ratio
l2 = -advantages * ratio.clamp(1 - clip_low, 1 + clip_high)
pg = torch.maximum(l1, l2)
mask = mask.to(pg.dtype)
return (pg * mask).sum() / mask.sum().clamp(min=1.0)
def clip_fraction(logp, old_logp, clip):
r = (logp - old_logp).exp()
clipped = (r > 1 + clip) | (r < 1 - clip)
return clipped.float().mean()
def value_loss(values, old_values, returns, clip):
if clip is None:
return 0.5 * ((values - returns) ** 2).mean()
vc = old_values + (values - old_values).clamp(-clip, clip)
l1 = (values - returns) ** 2
l2 = (vc - returns) ** 2
return 0.5 * torch.maximum(l1, l2).mean()
def huber_value_loss(values, returns, delta):
return F.huber_loss(values, returns, delta=delta)
# ---------------------------------------------------------------------------
# Whitening / normalization
# ---------------------------------------------------------------------------
def whiten(values, mask, shift_mean):
if mask is None:
mask = torch.ones_like(values)
mask = mask.to(values.dtype)
mean = (values * mask).sum() / mask.sum().clamp(min=1.0)
var = ((values - mean) ** 2 * mask).sum() / mask.sum().clamp(min=1.0)
whitened = (values - mean) * torch.rsqrt(var + 1e-8)
if not shift_mean:
whitened = whitened + mean
return whitened
def masked_whiten(values, mask, shift_mean):
return whiten(values, mask, shift_mean) * mask.to(values.dtype)
def normalize(x, eps):
return (x - x.mean()) / (x.std(unbiased=True) + eps)
def advantage_mean_std(advantages, mask):
mask = mask.to(advantages.dtype)
n = mask.sum().clamp(min=1.0)
mean = (advantages * mask).sum() / n
var = ((advantages - mean) ** 2 * mask).sum() / n
return torch.stack([mean, var.sqrt()])
# ---------------------------------------------------------------------------
# Losses over logits
# ---------------------------------------------------------------------------
def smoothed_nll(logits, labels, smoothing):
lp = torch.log_softmax(logits, dim=-1)
nll = -lp.gather(-1, labels.unsqueeze(-1)).squeeze(-1)
smooth = -lp.mean(-1)
loss = (1 - smoothing) * nll + smoothing * smooth
return loss.mean()
def cross_entropy(logits, labels, ignore_index):
return F.cross_entropy(
logits.reshape(-1, logits.shape[-1]),
labels.reshape(-1),
ignore_index=ignore_index,
)
# ---------------------------------------------------------------------------
# Sampling masks / selection
# ---------------------------------------------------------------------------
def top_p_mask(probs, p):
sorted_probs, sorted_idx = torch.sort(probs, descending=True, dim=-1)
cumsum = sorted_probs.cumsum(-1)
keep = (cumsum - sorted_probs) < p
kept = sorted_probs * keep
kept = kept / kept.sum(-1, keepdim=True)
out = torch.zeros_like(probs)
out.scatter_(-1, sorted_idx, kept)
return out
def top_k_mask(logits, k):
idx = logits.topk(k, dim=-1).indices
mask = torch.zeros_like(logits, dtype=torch.bool)
mask.scatter_(-1, idx, True)
return mask
def argmax_tokens(logits):
return logits.argmax(-1)
# ---------------------------------------------------------------------------
# Statistics
# ---------------------------------------------------------------------------
def mode_label(labels):
return torch.mode(labels, dim=-1).values
def median_reward(rewards):
return torch.quantile(rewards, 0.5)
def masked_max(values, mask):
mb = mask.to(torch.bool)
if not bool(mb.any()):
return torch.zeros((), dtype=values.dtype)
return values[mb].max()
def quantile_lower(x, q):
return torch.quantile(x, q, interpolation='lower')
# ---------------------------------------------------------------------------
# Misc tensor ops
# ---------------------------------------------------------------------------
def pad_mask_from_lengths(lengths, max_len):
ar = torch.arange(max_len, device=lengths.device)
return (ar.unsqueeze(0) < lengths.unsqueeze(1)).float()
def first_nonzero_index(mask):
nz = mask != 0
idx = nz.to(torch.int).argmax(-1)
any_nz = nz.any(-1)
return torch.where(any_nz, idx, torch.full_like(idx, -1))
def cumulative_max(x):
cm = torch.cummax(x, dim=0).values
out = torch.empty_like(x)
out[0] = x[0]
if x.shape[0] > 1:
out[1:] = cm[:-1]
return out
def bucketize_reward(rewards, boundaries):
return torch.bucketize(rewards, boundaries, right=True)
# ---------------------------------------------------------------------------
# Composite RL objectives
# ---------------------------------------------------------------------------
def grpo_objective(logits, old_logits, ref_logits, labels, completion_mask,
rewards, group_size, beta, clip_low, clip_high, scale_by_std,
kl_estimator):
lp = token_logprobs(logits, labels)
olp = token_logprobs(old_logits, labels)
rlp = token_logprobs(ref_logits, labels)
adv = grpo_advantages(rewards, group_size, scale_by_std).unsqueeze(1)
ratio = (lp - olp).exp()
l1 = -adv * ratio
l2 = -adv * ratio.clamp(1 - clip_low, 1 + clip_high)
pg = torch.maximum(l1, l2)
kl = kl_penalty(lp, rlp, kl_estimator)
per = pg + beta * kl
mask = completion_mask.to(per.dtype)
return (per * mask).sum() / mask.sum().clamp(min=1.0)
def rloo_objective(logits, old_logits, labels, mask, rewards, group_size,
clip_low, clip_high):
lp = token_logprobs(logits, labels)
olp = token_logprobs(old_logits, labels)
adv = rloo_advantages(rewards, group_size).unsqueeze(1)
ratio = (lp - olp).exp()
l1 = -adv * ratio
l2 = -adv * ratio.clamp(1 - clip_low, 1 + clip_high)
pg = torch.maximum(l1, l2)
m = mask.to(pg.dtype)
return (pg * m).sum() / m.sum().clamp(min=1.0)
def reverse_kl_objective(logits, ref_logits, labels, mask, advantages, beta):
lp = token_logprobs(logits, labels)
rlp = token_logprobs(ref_logits, labels)
kl = reverse_kl(lp, rlp)
adv = advantages.unsqueeze(1)
per = -adv * lp + beta * kl
m = mask.to(per.dtype)
return (per * m).sum() / m.sum().clamp(min=1.0)
def ppo_objective(rewards, values, old_values, logp, old_logp, next_value,
gamma, lam, clip_low, clip_high, vf_clip, vf_coef):
adv = gae(rewards, values, next_value, gamma, lam)
returns = adv + values
mean = adv.mean()
var = ((adv - mean) ** 2).mean()
advw = (adv - mean) * torch.rsqrt(var + 1e-8)
ratio = (logp - old_logp).exp()
l1 = -advw * ratio
l2 = -advw * ratio.clamp(1 - clip_low, 1 + clip_high)
pg = torch.maximum(l1, l2).mean()
vl = value_loss(values, old_values, returns, vf_clip)
return pg + vf_coef * vl
Full verification suite
python3 -c "
import torch, importlib.util, reforacle
spec=importlib.util.spec_from_file_location('obj','/app/objectives.py')
obj=importlib.util.module_from_spec(spec); spec.loader.exec_module(obj)
def close(a,b):
a=torch.as_tensor(a); b=torch.as_tensor(b)
return torch.allclose(a.float(),b.float(),atol=1e-5,rtol=1e-4)
fails=[]
def check(name, fn_args_gen, n=8):
for i in range(n):
g=torch.Generator().manual_seed(1000+i)
try:
args=fn_args_gen(g)
r=getattr(reforacle,name)(*args)
m=getattr(obj,name)(*args)
except Exception as e:
fails.append((name,i,'EXC',str(e)[:80])); return
if not close(r,m):
fails.append((name,i,r,m)); return
print('OK',name)
rn=lambda g,*s: torch.randn(*s,generator=g)
mk=lambda g,*s: (torch.rand(*s,generator=g)>0.4).float()
check('masked_mean', lambda g:(rn(g,3,5),mk(g,3,5)))
check('masked_mean', lambda g:(rn(g,3,5),mk(g,3,5),1))
check('masked_mean', lambda g:(rn(g,3,5),mk(g,3,5),0))
check('masked_sum', lambda g:(rn(g,3,5),mk(g,3,5),1))
check('masked_sum', lambda g:(rn(g,3,5),mk(g,3,5)))
check('logsumexp', lambda g:(rn(g,3,5),1))
check('log_softmax', lambda g:(rn(g,3,5),1))
check('token_logprobs', lambda g:(rn(g,2,4,6),torch.randint(0,6,(2,4),generator=g)))
check('selective_logprobs', lambda g:(rn(g,2,4,6),torch.randint(0,6,(2,4),generator=g),mk(g,2,4)))
check('sequence_logprob', lambda g:(rn(g,2,4,6),torch.randint(0,6,(2,4),generator=g),mk(g,2,4),True))
check('sequence_logprob', lambda g:(rn(g,2,4,6),torch.randint(0,6,(2,4),generator=g),mk(g,2,4),False))
check('logprob_at_temperature', lambda g:(rn(g,2,4,6),torch.randint(0,6,(2,4),generator=g),1.5))
check('entropy', lambda g:(rn(g,2,4,6),mk(g,2,4)))
check('normalized_entropy', lambda g:(rn(g,2,4,6),mk(g,2,4)))
check('dpo_loss', lambda g:(rn(g,5),rn(g,5),rn(g,5),rn(g,5),0.1,0.15))
check('ipo_loss', lambda g:(rn(g,5),rn(g,5),rn(g,5),rn(g,5),0.2))
check('bradley_terry_logit', lambda g:(rn(g,5),rn(g,5),0.3))
check('grpo_advantages', lambda g:(rn(g,6),3,True))
check('grpo_advantages', lambda g:(rn(g,6),3,False))
check('rloo_advantages', lambda g:(rn(g,6),3))
check('group_mean_baseline', lambda g:(rn(g,6),2))
check('gae', lambda g:(rn(g,5),rn(g,5),torch.tensor(0.3),0.99,0.95))
check('lambda_returns', lambda g:(rn(g,5),rn(g,5),torch.tensor(0.3),0.99,0.95))
check('discounted_returns', lambda g:(rn(g,5),0.97))
for est in ['k1','k2','k3']:
check('kl_penalty', lambda g,e=est:(rn(g,6),rn(g,6),e))
check('reverse_kl', lambda g:(rn(g,6),rn(g,6)))
check('symmetric_kl', lambda g:(rn(g,6),rn(g,6)))
check('importance_ratio', lambda g:(rn(g,6),rn(g,6),0.2))
check('importance_ratio', lambda g:(rn(g,6),rn(g,6),None))
check('clipped_pg_loss', lambda g:(rn(g,2,4),rn(g,2,4),rn(g,2,4),mk(g,2,4),0.2,0.3))
check('clip_fraction', lambda g:(rn(g,6),rn(g,6),0.2))
check('value_loss', lambda g:(rn(g,6),rn(g,6),rn(g,6),0.2))
check('huber_value_loss', lambda g:(rn(g,6),rn(g,6),1.0))
check('huber_value_loss', lambda g:(rn(g,6),rn(g,6),0.5))
check('whiten', lambda g:(rn(g,6),mk(g,6),True))
check('whiten', lambda g:(rn(g,6),mk(g,6),False))
check('masked_whiten', lambda g:(rn(g,6),mk(g,6),True))
check('masked_whiten', lambda g:(rn(g,6),mk(g,6),False))
check('normalize', lambda g:(rn(g,6),1e-8))
check('advantage_mean_std', lambda g:(rn(g,2,4),mk(g,2,4)))
check('smoothed_nll', lambda g:(rn(g,2,4,6),torch.randint(0,6,(2,4),generator=g),0.1))
check('cross_entropy', lambda g:(rn(g,2,4,6),torch.randint(0,6,(2,4),generator=g),-100))
check('top_p_mask', lambda g:(torch.softmax(rn(g,3,5),-1),0.8))
check('top_k_mask', lambda g:(rn(g,3,5),2))
check('argmax_tokens', lambda g:(rn(g,3,5),))
check('mode_label', lambda g:(torch.randint(0,4,(7,),generator=g),))
check('median_reward', lambda g:(rn(g,7),))
check('masked_max', lambda g:(rn(g,6),mk(g,6)))
check('quantile_lower', lambda g:(rn(g,7),0.25))
check('pad_mask_from_lengths', lambda g:(torch.randint(1,5,(3,),generator=g),5))
check('first_nonzero_index', lambda g:(mk(g,3,5),))
check('cumulative_max', lambda g:(rn(g,6),))
check('bucketize_reward', lambda g:(rn(g,6),torch.tensor([-0.5,0.0,0.5])))
check('dpo_sequence_loss', lambda g:(rn(g,2,4,6),rn(g,2,4,6),rn(g,2,4,6),rn(g,2,4,6),torch.randint(0,6,(2,4),generator=g),torch.randint(0,6,(2,4),generator=g),mk(g,2,4),mk(g,2,4),0.1,0.1))
check('grpo_objective', lambda g:(rn(g,4,3,5),rn(g,4,3,5),rn(g,4,3,5),torch.randint(0,5,(4,3),generator=g),mk(g,4,3),rn(g,4),2,0.04,0.2,0.2,True,'k3'))
check('grpo_objective', lambda g:(rn(g,4,3,5),rn(g,4,3,5),rn(g,4,3,5),torch.randint(0,5,(4,3),generator=g),mk(g,4,3),rn(g,4),2,0.04,0.2,0.2,False,'k1'))
check('rloo_objective', lambda g:(rn(g,4,3,5),rn(g,4,3,5),torch.randint(0,5,(4,3),generator=g),mk(g,4,3),rn(g,4),2,0.2,0.2))
check('reverse_kl_objective', lambda g:(rn(g,4,3,5),rn(g,4,3,5),torch.randint(0,5,(4,3),generator=g),mk(g,4,3),rn(g,4),0.1))
check('ppo_objective', lambda g:(rn(g,5),rn(g,5),rn(g,5),rn(g,5),rn(g,5),torch.tensor(0.3),0.99,0.95,0.2,0.2,0.2,0.5))
print()
print('FAILS:', len(fails))
for f in fails: print(f)
"
OK masked_mean
OK masked_mean
OK masked_mean
OK masked_sum
OK masked_sum
OK logsumexp
OK log_softmax
OK token_logprobs
OK selective_logprobs
OK sequence_logprob
OK sequence_logprob
OK logprob_at_temperature
OK entropy
OK normalized_entropy
OK dpo_loss
OK ipo_loss
OK bradley_terry_logit
OK grpo_advantages
OK grpo_advantages
OK rloo_advantages
OK group_mean_baseline
OK gae
OK lambda_returns
OK discounted_returns
OK kl_penalty
OK kl_penalty
OK kl_penalty
OK reverse_kl
OK symmetric_kl
OK importance_ratio
OK importance_ratio
OK clipped_pg_loss
OK clip_fraction
OK value_loss
OK huber_value_loss
OK huber_value_loss
OK whiten
OK whiten
OK masked_whiten
OK normalize
OK advantage_mean_std
OK smoothed_nll
OK cross_entropy
OK top_p_mask
OK top_k_mask
OK argmax_tokens
OK median_reward
OK masked_max
OK quantile_lower
OK pad_mask_from_lengths
OK first_nonzero_index
OK cumulative_max
OK bucketize_reward
OK dpo_sequence_loss
OK grpo_objective
OK grpo_objective
OK rloo_objective
OK reverse_kl_objective
OK ppo_objective
FAILS: 2
('masked_whiten', 0, tensor([-1.2973, 0.0939, 0.5693, 0.0939, 1.0097, 0.0939]), tensor([-1.2973, -0.0000, 0.5693, 0.0000, 1.0097, -0.0000]))
('mode_label', 6, tensor(3), tensor(1))
[stdout]
OK masked_mean
OK masked_mean
OK masked_mean
OK masked_sum
OK masked_sum
OK logsumexp
OK log_softmax
OK token_logprobs
OK selective_logprobs
OK sequence_logprob
OK sequence_logprob
OK logprob_at_temperature
OK entropy
OK normalized_entropy
OK dpo_loss
OK ipo_loss
OK bradley_terry_logit
OK grpo_advantages
OK grpo_advantages
OK rloo_advantages
OK group_mean_baseline
OK gae
OK lambda_returns
OK discounted_returns
OK kl_penalty
OK kl_penalty
OK kl_penalty
OK reverse_kl
OK symmetric_kl
OK importance_ratio
OK importance_ratio
OK clipped_pg_loss
OK clip_fraction
OK value_loss
OK huber_value_loss
OK huber_value_loss
OK whiten
OK whiten
OK masked_whiten
OK normalize
OK advantage_mean_std
OK smoothed_nll
OK cross_entropy
OK top_p_mask
OK top_k_mask
OK argmax_tokens
OK median_reward
OK masked_max
OK quantile_lower
OK pad_mask_from_lengths
OK first_nonzero_index
OK cumulative_max
OK bucketize_reward
OK dpo_sequence_loss
OK grpo_objective
OK grpo_objective
OK rloo_objective
OK reverse_kl_objective
OK ppo_objective
FAILS: 2
('masked_whiten', 0, tensor([-1.2973, 0.0939, 0.5693, 0.0939, 1.0097, 0.0939]), tensor([-1.2973, -0.0000, 0.5693, 0.0000, 1.0097, -0.0000]))
('mode_label', 6, tensor(3), tensor(1))masked_whiten reexamine
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(1000)
v=torch.randn(6,generator=g); m=(torch.rand(6,generator=g)>0.4).float()
print('v',v); print('m',m)
print('mw T', reforacle.masked_whiten(v,m,True))
print('mw F', reforacle.masked_whiten(v,m,False))
print('whiten T', reforacle.whiten(v,m,True))
print('whiten F', reforacle.whiten(v,m,False))
mean=(v*m).sum()/m.sum(); var=((v-mean)**2*m).sum()/m.sum()
print('mean',mean.item(),'std',var.sqrt().item())
# masked positions value 0.0939 vs?
print('(v-mean)/std at masked', ((v-mean)/torch.sqrt(var+1e-8)))
"
v tensor([-1.1720, -0.3929, 0.5265, 1.1065, 0.9273, -1.7421]) m tensor([1., 0., 1., 0., 1., 0.]) mw T tensor([-1.3912, -0.0000, 0.4754, 0.0000, 0.9158, -0.0000]) mw F tensor([-1.2973, 0.0939, 0.5693, 0.0939, 1.0097, 0.0939]) whiten T tensor([-1.3912, -0.5349, 0.4754, 1.1128, 0.9158, -2.0176]) whiten F tensor([-1.2973, -0.4410, 0.5693, 1.2067, 1.0097, -1.9237]) mean 0.09390169382095337 std 0.9099937677383423 (v-mean)/std at masked tensor([-1.3912, -0.5349, 0.4754, 1.1128, 0.9158, -2.0176]) [stdout] v tensor([-1.1720, -0.3929, 0.5265, 1.1065, 0.9273, -1.7421]) m tensor([1., 0., 1., 0., 1., 0.]) mw T tensor([-1.3912, -0.0000, 0.4754, 0.0000, 0.9158, -0.0000]) mw F tensor([-1.2973, 0.0939, 0.5693, 0.0939, 1.0097, 0.0939]) whiten T tensor([-1.3912, -0.5349, 0.4754, 1.1128, 0.9158, -2.0176]) whiten F tensor([-1.2973, -0.4410, 0.5693, 1.2067, 1.0097, -1.9237]) mean 0.09390169382095337 std 0.9099937677383423 (v-mean)/std at masked tensor([-1.3912, -0.5349, 0.4754, 1.1128, 0.9158, -2.0176])
mode tie behavior
python3 -c "
import torch, reforacle
g=torch.Generator().manual_seed(1006)
lab=torch.randint(0,4,(7,),generator=g)
print('lab',lab, 'counts',torch.bincount(lab))
print('ref',reforacle.mode_label(lab))
print('torch.mode',torch.mode(lab,-1).values)
print('bincount argmax',torch.bincount(lab).argmax())
# test more ties
for s in range(1000,1010):
g=torch.Generator().manual_seed(s); lab=torch.randint(0,5,(6,),generator=g)
print(lab.tolist(),'ref',reforacle.mode_label(lab).item(),'bincount',torch.bincount(lab).argmax().item(),'mode',torch.mode(lab,-1).values.item())
"
lab tensor([0, 1, 3, 2, 3, 1, 2]) counts tensor([1, 2, 2, 2]) ref tensor(3) torch.mode tensor(1) bincount argmax tensor(1) [2, 4, 2, 3, 1, 4] ref 4 bincount 2 mode 2 [2, 2, 2, 0, 1, 1] ref 2 bincount 2 mode 2 [2, 0, 2, 1, 4, 0] ref 2 bincount 0 mode 0 [1, 4, 4, 1, 1, 2] ref 1 bincount 1 mode 1 [2, 2, 2, 1, 0, 1] ref 2 bincount 2 mode 2 [4, 0, 3, 3, 3, 3] ref 3 bincount 3 mode 3 [0, 3, 1, 1, 1, 4] ref 1 bincount 1 mode 1 [1, 0, 0, 0, 3, 4] ref 0 bincount 0 mode 0 [2, 4, 0, 1, 0, 3] ref 0 bincount 0 mode 0 [3, 4, 2, 2, 4, 1] ref 4 bincount 2 mode 2 [stdout] lab tensor([0, 1, 3, 2, 3, 1, 2]) counts tensor([1, 2, 2, 2]) ref tensor(3) torch.mode tensor(1) bincount argmax tensor(1) [2, 4, 2, 3, 1, 4] ref 4 bincount 2 mode 2 [2, 2, 2, 0, 1, 1] ref 2 bincount 2 mode 2 [2, 0, 2, 1, 4, 0] ref 2 bincount 0 mode 0 [1, 4, 4, 1, 1, 2] ref 1 bincount 1 mode 1 [2, 2, 2, 1, 0, 1] ref 2 bincount 2 mode 2 [4, 0, 3, 3, 3, 3] ref 3 bincount 3 mode 3 [0, 3, 1, 1, 1, 4] ref 1 bincount 1 mode 1 [1, 0, 0, 0, 3, 4] ref 0 bincount 0 mode 0 [2, 4, 0, 1, 0, 3] ref 0 bincount 0 mode 0 [3, 4, 2, 2, 4, 1] ref 4 bincount 2 mode 2
/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 = (values * mask).sum() / mask.sum().clamp(min=1.0)
var = ((values - mean) ** 2 * mask).sum() / mask.sum().clamp(min=1.0)
whitened = (values - mean) * torch.rsqrt(var + 1e-8)
whitened = whitened * mask
if not shift_mean:
whitened = whitened + mean
return whitened/app/objectives.py
def mode_label(labels):
return torch.mode(labels, dim=-1).valuesdef mode_label(labels):
eq = labels.unsqueeze(-1) == labels.unsqueeze(-2)
counts = eq.sum(-1)
maxc = counts.max(-1, keepdim=True).values
is_mode = counts == maxc
fill = torch.full_like(labels, torch.iinfo(labels.dtype).min)
cand = torch.where(is_mode, labels, fill)
return cand.max(-1).valuesRe-verify fixes
python3 -c "
import torch, importlib.util, reforacle
spec=importlib.util.spec_from_file_location('obj','/app/objectives.py')
obj=importlib.util.module_from_spec(spec); spec.loader.exec_module(obj)
def close(a,b):
a=torch.as_tensor(a); b=torch.as_tensor(b); return torch.allclose(a.float(),b.float(),atol=1e-5,rtol=1e-4)
fails=[]
def check(name, gen, n=20):
for i in range(n):
g=torch.Generator().manual_seed(3000+i)
try:
args=gen(g); r=getattr(reforacle,name)(*args); m=getattr(obj,name)(*args)
except Exception as e: fails.append((name,i,'EXC',str(e)[:80])); return
if not close(r,m): fails.append((name,i,r,m)); return
print('OK',name)
rn=lambda g,*s: torch.randn(*s,generator=g)
mk=lambda g,*s: (torch.rand(*s,generator=g)>0.4).float()
check('masked_whiten', lambda g:(rn(g,6),mk(g,6),True))
check('masked_whiten', lambda g:(rn(g,6),mk(g,6),False))
check('mode_label', lambda g:(torch.randint(0,4,(7,),generator=g),), n=50)
check('mode_label', lambda g:(torch.randint(0,3,(6,),generator=g),), n=50)
check('mode_label', lambda g:(torch.randint(0,5,(4,6),generator=g),), n=30)
print('FAILS',len(fails))
for f in fails: print(f)
"
OK masked_whiten OK masked_whiten OK mode_label OK mode_label OK mode_label FAILS 0 [stdout] OK masked_whiten OK masked_whiten OK mode_label OK mode_label OK mode_label FAILS 0
Stress test varied hyperparams/shapes
python3 -c "
import torch, importlib.util, reforacle, random
spec=importlib.util.spec_from_file_location('obj','/app/objectives.py')
obj=importlib.util.module_from_spec(spec); spec.loader.exec_module(obj)
def close(a,b):
a=torch.as_tensor(a); b=torch.as_tensor(b)
if a.shape!=b.shape: return False
return torch.allclose(a.float(),b.float(),atol=1e-5,rtol=1e-4)
fails=[]
def check(name, gen, n=15):
for i in range(n):
g=torch.Generator().manual_seed(7000+i*7)
try:
args=gen(g,i); r=getattr(reforacle,name)(*args); m=getattr(obj,name)(*args)
except Exception as e: fails.append((name,i,'EXC',str(e)[:100])); return
if not close(r,m): fails.append((name,i,r,m)); return
print('OK',name)
rn=lambda g,*s: torch.randn(*s,generator=g)
mk=lambda g,*s: (torch.rand(*s,generator=g)>0.4).float()
ri=lambda g,hi,*s: torch.randint(0,hi,s,generator=g)
# varied hyperparams
check('dpo_loss', lambda g,i:(rn(g,4),rn(g,4),rn(g,4),rn(g,4),0.05+0.1*i,0.05*i%0.4))
check('ipo_loss', lambda g,i:(rn(g,4),rn(g,4),rn(g,4),rn(g,4),0.1+0.1*i))
check('grpo_advantages', lambda g,i:(rn(g,12),[2,3,4,6][i%4],bool(i%2)))
check('rloo_advantages', lambda g,i:(rn(g,12),[2,3,4,6][i%4]))
check('gae', lambda g,i:(rn(g,3+i),rn(g,3+i),torch.tensor(float(rn(g,1))),0.9+0.01*i,0.9))
check('discounted_returns', lambda g,i:(rn(g,3+i),0.9+0.009*i))
check('value_loss', lambda g,i:(rn(g,6),rn(g,6),rn(g,6),0.1+0.1*i))
check('huber_value_loss', lambda g,i:(rn(g,6),rn(g,6),0.3+0.2*i))
check('cross_entropy', lambda g,i:(rn(g,3,7),ri(g,7,3),-100))
check('smoothed_nll', lambda g,i:(rn(g,3,7),ri(g,7,3),0.05*i))
check('top_p_mask', lambda g,i:(torch.softmax(rn(g,4,8),-1),0.5+0.03*i))
check('top_k_mask', lambda g,i:(rn(g,4,8),1+i%5))
check('bucketize_reward', lambda g,i:(rn(g,8),torch.sort(rn(g,4))[0]))
check('quantile_lower', lambda g,i:(rn(g,10),0.1*(i%9)+0.05))
check('clipped_pg_loss', lambda g,i:(rn(g,3,5),rn(g,3,5),rn(g,3,5),mk(g,3,5),0.1+0.05*i,0.1+0.03*i))
check('clip_fraction', lambda g,i:(rn(g,10),rn(g,10),0.1+0.05*i))
check('importance_ratio', lambda g,i:(rn(g,8),rn(g,8),None if i%2 else 0.1+0.05*i))
check('grpo_objective', lambda g,i:(rn(g,6,3,5),rn(g,6,3,5),rn(g,6,3,5),ri(g,5,6,3),mk(g,6,3),rn(g,6),[2,3,6][i%3],0.01+0.02*i,0.1+0.05*i,0.1+0.05*i,bool(i%2),['k1','k2','k3'][i%3]))
check('rloo_objective', lambda g,i:(rn(g,6,3,5),rn(g,6,3,5),ri(g,5,6,3),mk(g,6,3),rn(g,6),[2,3,6][i%3],0.1+0.02*i,0.1+0.02*i))
check('reverse_kl_objective', lambda g,i:(rn(g,5,4,6),rn(g,5,4,6),ri(g,6,5,4),mk(g,5,4),rn(g,5),0.05+0.05*i))
check('ppo_objective', lambda g,i:(rn(g,4+i),rn(g,4+i),rn(g,4+i),rn(g,4+i),rn(g,4+i),torch.tensor(float(rn(g,1))),0.95,0.9,0.1+0.02*i,0.1+0.02*i,0.1+0.02*i,0.5))
check('dpo_sequence_loss', lambda g,i:(rn(g,3,4,6),rn(g,3,4,6),rn(g,3,4,6),rn(g,3,4,6),ri(g,6,3,4),ri(g,6,3,4),mk(g,3,4),mk(g,3,4),0.05+0.05*i,0.05*i%0.3))
check('masked_mean', lambda g,i:(rn(g,4,5),mk(g,4,5),[None,0,1][i%3]))
check('advantage_mean_std', lambda g,i:(rn(g,3,4),mk(g,3,4)))
check('sequence_logprob', lambda g,i:(rn(g,3,5,7),ri(g,7,3,5),mk(g,3,5),bool(i%2)))
check('first_nonzero_index', lambda g,i:(mk(g,4,6),))
check('cumulative_max', lambda g,i:(rn(g,7),) if i%2 else (rn(g,5,3),))
check('pad_mask_from_lengths', lambda g,i:(torch.randint(0,6,(4,),generator=g),6))
print('FAILS',len(fails))
for f in fails: print(f)
"
OK dpo_loss OK ipo_loss OK grpo_advantages OK rloo_advantages OK gae OK discounted_returns OK value_loss OK huber_value_loss OK cross_entropy OK smoothed_nll OK top_p_mask OK top_k_mask OK bucketize_reward OK quantile_lower OK clipped_pg_loss OK clip_fraction OK importance_ratio OK grpo_objective OK rloo_objective OK reverse_kl_objective OK ppo_objective OK dpo_sequence_loss OK masked_mean OK advantage_mean_std OK sequence_logprob OK first_nonzero_index OK cumulative_max OK pad_mask_from_lengths FAILS 0 [stdout] OK dpo_loss OK ipo_loss OK grpo_advantages OK rloo_advantages OK gae OK discounted_returns OK value_loss OK huber_value_loss OK cross_entropy OK smoothed_nll OK top_p_mask OK top_k_mask OK bucketize_reward OK quantile_lower OK clipped_pg_loss OK clip_fraction OK importance_ratio OK grpo_objective OK rloo_objective OK reverse_kl_objective OK ppo_objective OK dpo_sequence_loss OK masked_mean OK advantage_mean_std OK sequence_logprob OK first_nonzero_index OK cumulative_max OK pad_mask_from_lengths FAILS 0
Confirm standalone, no reforacle dependency
grep -n "reforacle\|import" /app/objectives.py; echo "---"; python3 -c "
import sys
# simulate clean sandbox: block reforacle
class Blk:
def find_spec(self,n,p,t=None):
if n=='reforacle': raise ImportError('blocked')
sys.meta_path.insert(0,Blk())
import importlib.util
spec=importlib.util.spec_from_file_location('obj','/app/objectives.py')
obj=importlib.util.module_from_spec(spec); spec.loader.exec_module(obj)
import torch
print('loads cleanly; masked_mean sanity:', obj.masked_mean(torch.tensor([1.,2,3]),torch.tensor([1.,0,1])).item())
"1:from __future__ import annotations 3:import torch 4:import torch.nn.functional as F 205:def importance_ratio(logp, old_logp, clip): --- loads cleanly; masked_mean sanity: 2.0 [stdout] 1:from __future__ import annotations 3:import torch 4:import torch.nn.functional as F 205:def importance_ratio(logp, old_logp, clip): --- loads cleanly; masked_mean sanity: 2.0
No file changes, the agent left the workspace untouched.
=== SyncVals verifier (post-training objectives, sealed oracle) === FUNCTIONS PASSED: 46/49 (partial_score=0.9388) 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_983eb99d2eaf42ca. Re-running the agent live requires EVAL_PLATFORM_ENABLE_OAUTH_SMOKE=1 and is non-deterministic.
Trial trial_983eb99d2eaf42ca · verifier authoritative; classifier explanatory.