SyncValsverifier → artifact → classifier → verdict
SyncVals · Trajectory

reverse-engineer-objectives

claude-code claude-opus-4-8 ✗ failed GOOD_FAILURE ↑ View task
Solved from the instruction alone, tests/ and solution/ were withheld from the agent's workspace and restored only for grading.
Reward = tests/test.sh exit code (0 → resolved); the classification below is post-hoc and cannot change it.
Classification , post-hoc; cannot change the reward
GOOD_FAILUREHonest miss, the agent ran correctly but couldn't solve it. Expected for a hard task; the task is sound.
SubtypeComplexity Overwhelm
EvidenceAgent passed 45/49 functions locally (verified via 25-trial validation script showing 'ALL PASS') but verifier shows only 45/49 match on held-out inputs, requiring 49/49 for task success. Test output: 'FUNCTIONS PASSED: 45/49 (partial_score=0.9184) REWARD = 0'. Agent trajectory shows systematic probing discovering conventions (grpo unbiased=False std, mode tie-breaking returns max, bucketize right=True, masked_mean empty mask→0.0, symmetric_kl formula, etc.) but 4 functions diverged on grading's unseen test cases.
Root causeThe task explicitly requires exact numerical agreement on held-out inputs the agent never sees. Despite thorough reverse-engineering and 92% local accuracy, 4 of 49 functions failed on grading's hidden test inputs, likely due to untested edge cases or subtle numerical precision issues (the tolerance is tight: atol=1e-6/rtol=1e-5 for primitives, atol=1e-5/rtol=1e-4 for composites).
RecommendationN/A - task is fine. This is expected behavior for a reverse-engineering challenge with hidden grading inputs. The agent demonstrated solid understanding but couldn't exhaustively cover all edge cases without access to the test set.
Trajectory
Tool-by-tool agent trajectory
67 tool calls · 4 tool types · 69 steps
Reimplement the post-training objectives library in `/app/objectives.py` so that every function reproduces the sealed reference module `reforacle` exactly. `reforacle` is importable and callable (`import reforacle`) but its source cannot be read. It is the only ground truth; there is no written specification of what any function computes or of the conventions it follows. Recover its behavior by calling it. `/app/objectives.py` contains forty-nine function stubs. Keep every signature exactly as given and implement every function. The names indicate the standard post-training quantity each is meant to produce, but the reference's exact definition is whatever `reforacle` returns , not what you assume. Determine each definition by probing. Grading compares your implementation to `reforacle` on held-out inputs you do not see, function by function, requiring exact numerical agreement (primitives within `atol=1e-6, rtol=1e-5`; composite objectives within `atol=1e-5, rtol=1e-4`). Each function passes only if it agrees on every graded input, and the task passes only if all forty-nine functions pass. Agreement on the inputs you happen to try does not imply agreement on the inputs you are graded on. Grading runs your `/app/objectives.py` in a clean sandbox that does not contain `reforacle`, so any submission that imports, wraps, calls, loads, or otherwise depends on the reference at run time fails outright. Your file must be a standalone implementation whose only dependency is `torch`. Do not read, decompile, disassemble, or otherwise extract the compiled reference; recover its behavior only by calling it and comparing. `python3 /app/probe.py` shows the calling pattern; `bash /app/timer.sh` reports remaining budget.

/app/objectives.py

contents
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

contents
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

Check remaining budget

shell
bash /app/timer.sh
output
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

List reforacle functions

shell
python3 -c "import reforacle; print([x for x in dir(reforacle) if not x.startswith('_')])"
output
['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']
write · /app/p1.py
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

shell
python3 p1.py
output
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
write · /app/p2.py
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

shell
python3 p2.py
output
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]
write · /app/p3.py
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

shell
python3 p3.py
output
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]
write · /app/p4.py
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

shell
python3 p4.py
output
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)
write · /app/p5.py
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

shell
python3 p5.py
output
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
write · /app/p6.py
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

shell
python3 p6.py
output
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 failure
write · /app/p7.py
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()

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

shell
python3 p7.py
output
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]
write · /app/p8.py
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

shell
python3 p8.py
output
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])
write · /app/p9.py
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

shell
python3 p9.py
output
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])
write · /app/p10.py
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

shell
python3 p10.py
output
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)
write · /app/p11.py
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

shell
python3 p11.py
output
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
write · /app/p12.py
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

shell
python3 p12.py
output
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
write · /app/p13.py
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

shell
python3 p13.py
output
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
write · /app/p14.py
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

shell
python3 p14.py
output
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 []
write · /app/p15.py
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

shell
python3 p15.py
output
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 failure
write · /app/p16.py
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; 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

shell
python3 p16.py
output
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
write · /app/p17.py
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

shell
python3 p17.py
output
-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
write · /app/p18.py
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

shell
python3 p18.py
output
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]])
write · /app/p19.py
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,)

shell
python3 p19.py
output
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
write · /app/p20.py
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

shell
python3 p20.py
output
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
write · /app/p21.py
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

shell
python3 p21.py
output
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
write · /app/objectives.py
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)
write · /app/validate.py
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

shell
python3 validate.py 2>&1 | tail -40
output
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']
write · /app/p22.py
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

shell
python3 p22.py
output
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

− removed
def masked_whiten(values, mask, shift_mean):
    return whiten(values, mask, shift_mean)
+ added
def masked_whiten(values, mask, shift_mean):
    return whiten(values, mask, shift_mean) * mask.to(values.dtype)

/app/objectives.py

− removed
def median_reward(rewards):
    return torch.median(rewards)
+ added
def median_reward(rewards):
    return torch.quantile(rewards, 0.5)

/app/objectives.py

− removed
def cumulative_max(x):
    return torch.cummax(x, dim=-1).values
+ added
def cumulative_max(x):
    c = torch.cummax(x, dim=0).values
    out = torch.empty_like(x)
    out[0] = x[0]
    out[1:] = c[:-1]
    return out

Re-run validation

shell
python3 validate.py 2>&1 | tail -15
output
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']
write · /app/p23.py
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

shell
python3 p23.py
output
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

− removed
def masked_whiten(values, mask, shift_mean):
    return whiten(values, mask, shift_mean) * mask.to(values.dtype)
+ added
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 out

Re-run validation

shell
python3 validate.py 2>&1 | tail -6
output
==== SUMMARY ====
tested 49 functions, 0 failing
ALL PASS

[stdout]
==== SUMMARY ====
tested 49 functions, 0 failing
ALL PASS
write · /app/p24.py
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

shell
python3 p24.py
output
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

shell
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())
"
output
[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

− removed
def mode_label(labels):
    return torch.mode(labels).values
+ added
def mode_label(labels):
    vals, counts = torch.unique(labels, return_counts=True)
    maxc = counts.max()
    return vals[counts == maxc].max()

/app/objectives.py

− removed
    return torch.bucketize(rewards, boundaries)
+ added
    return torch.bucketize(rewards, boundaries, right=True)

Verify mode/bucketize fixes and full validation

shell
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
output
[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

shell
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.py
output
imports: ['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

shell
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
output
Remaining time (hours:minutes):
1:43

[stdout]
Remaining time (hours:minutes):
1:43
All 49 functions reimplemented in `/app/objectives.py` and verified to match `reforacle` exactly across 25 randomized trials plus targeted edge cases. ## What I recovered by probing Every definition was determined by calling `reforacle` , never assumed. Key conventions I had to discover rather than guess: - **`grpo_advantages`**: group-wise centering, `std` with `unbiased=False` and `eps=1e-6`. - **`whiten` vs `masked_whiten`**: both use masked mean + *population* variance (`eps=1e-8`), but `masked_whiten` multiplies the whitened core by the mask *before* optionally re-adding the mean , `whiten` does not. - **`symmetric_kl` = `cosh(logp - ref_logp) - 1`** (equivalent to averaging the two k3 estimators). - **`kl_penalty`**: `k1 = logp - ref` (signed), `k2 = ½d²`, `k3 = exp(-d)-1+d`; `reverse_kl` equals the k3 form. - **`reverse_kl_objective`**: advantages are *per-sequence* `(N,)`, broadcast over tokens; loss = `masked_mean(-adv·logp + β·k3_token)`. - **`ppo_objective`**: GAE advantages, returns from *raw* (un-whitened) advantages, advantages whitened with population std for the PG term, value loss with the `0.5·max(...)` clipped form. - **`cumulative_max`**: *exclusive* cumulative max along **dim 0** (`out[i] = max(x[0..i-1])`, `out[0]=x[0]`). - **`median_reward`**: linear-interpolation quantile at 0.5 (not `torch.median`'s lower value). - **`mode_label`**: on ties, returns the **largest** value among max-count labels. - **`bucketize_reward`**: `right=True`. - **`masked_mean`**: empty mask → `0.0`. The file depends only on `torch` (+ stdlib `math`) , confirmed via AST scan that it never imports or references `reforacle`, so it runs standalone in the grading sandbox.
Final diff , what the agent changed

No file changes, the agent left the workspace untouched.

Verifier output , the proof behind the reward
tests/test.sh
=== 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.