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.
SubtypeIncomplete Implementation at Grading Time
EvidenceVerifier result shows 42/49 functions passing (reward=0.0, 85.7% score). Agent trajectory shows systematic reverse-engineering of all 49 functions, achieving 49/49 passing on local harness with multiple seeds and edge cases. Failed functions at grading time were: entropy, whiten, normalized_entropy. Agent trace shows these were later fixed (entropy/normalized_entropy harness bug at step 51, whiten variance bug at step 48), reaching 49/49 passing by step 54. Verifier output: 'FUNCTIONS PASSED: 42/49 (partial_score=0.8571) REWARD = 0 (1 iff ALL 49 functions match exactly)'
Root causeThe verifier ran against an intermediate/earlier version of the implementation (42/49 passing) before the agent completed all fixes. The agent later fixed the critical bugs (variance normalization in whiten, argument passing in entropy/normalized_entropy) and achieved 49/49 passing locally, but this final version was not what was tested by the verifier.
RecommendationN/A - task is fine. The task specification is clear, unambiguous, and achievable. The reference implementation is properly sealed. The grading logic is correct (all-or-nothing on exact match). The failure was timing: verifier executed before the agent's final fixes were complete. In a production setting, ensure the agent's final code submission is what gets graded, not an earlier snapshot.
Trajectory
Tool-by-tool agent trajectory
58 tool calls · 4 tool types · 60 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

Run probe to confirm reforacle works

shell
python3 probe.py
output
inputs:
  values = [1.5409960746765137, -0.293428897857666, -2.1787893772125244, 0.5684312582015991, -1.0845223665237427, -1.3985954523086548, 0.40334683656692505, 0.8380263447761536]
  mask   = [1.0, 1.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0]
reference masked_mean -> -0.07950486242771149
reference masked_mean (empty mask) -> 0.0

Probe any reference function this way to discover its exact numerical
convention, then implement it in /app/objectives.py to match.

[stdout]
inputs:
  values = [1.5409960746765137, -0.293428897857666, -2.1787893772125244, 0.5684312582015991, -1.0845223665237427, -1.3985954523086548, 0.40334683656692505, 0.8380263447761536]
  mask   = [1.0, 1.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0]
reference masked_mean -> -0.07950486242771149
reference masked_mean (empty mask) -> 0.0

Probe any reference function this way to discover its exact numerical
convention, then implement it in /app/objectives.py to match.

Probe primitives

shell
cat > /tmp/p1.py << 'EOF'
import torch, reforacle as R
g=torch.Generator().manual_seed(1)
def rn(*s): return torch.randn(*s, generator=g)

# masked_mean axis
v=rn(3,4); m=(torch.rand(3,4,generator=g)>0.4).float()
print("mm none", R.masked_mean(v,m))
print("mm ax0", R.masked_mean(v,m,axis=0))
print("mm ax1", R.masked_mean(v,m,axis=1))
print("ms none", R.masked_sum(v,m))
print("ms ax1", R.masked_sum(v,m,axis=1))
# logsumexp/log_softmax
x=rn(2,5)
print("lse", R.logsumexp(x,axis=1))
print("lsm", R.log_softmax(x,axis=1))
# token_logprobs
logits=rn(2,3,5); labels=torch.randint(0,5,(2,3),generator=g)
print("tlp", R.token_logprobs(logits,labels))
mask=(torch.rand(2,3,generator=g)>0.3).float()
print("slp", R.selective_logprobs(logits,labels,mask))
print("seqlp F", R.sequence_logprob(logits,labels,mask,False))
print("seqlp T", R.sequence_logprob(logits,labels,mask,True))
print("entropy", R.entropy(logits,mask))
print("normalize", R.normalize(rn(6),1e-8))
EOF
python3 /tmp/p1.py
output
Exit code 1
Traceback (most recent call last):
  File "/tmp/p1.py", line 1, in <module>
    import torch, reforacle as R
ModuleNotFoundError: No module named 'reforacle'

[error] tool reported failure

Run probe from /app

shell
cp /tmp/p1.py /app/_p1.py && python3 _p1.py
output
mm none tensor(-0.1978)
mm ax0 tensor([-1.0276, 0.0504, -0.7305, 0.5015])
mm ax1 tensor([ 0.3166, -0.4357, -1.0276])
ms none tensor(-1.3849)
ms ax1 tensor([ 0.9499, -1.3072, -1.0276])
lse tensor([3.7453, 2.9333])
lsm tensor([[-5.3180, -3.8686, -0.1584, -5.5766, -2.1466],
        [-4.2103, -2.6078, -3.4124, -1.5543, -0.4047]])
tlp tensor([[-2.6539, -1.9474, -2.8140],
        [-1.6476, -1.2624, -3.1213]])
slp tensor([-7.4154, -2.9100])
seqlp F tensor([-7.4154, -2.9100])
seqlp T tensor([-2.4718, -1.4550])
entropy tensor(1.3433)
normalize tensor([-1.8910, -0.0302, 0.8794, 0.6768, 0.4327, -0.0676])

[stdout]
mm none tensor(-0.1978)
mm ax0 tensor([-1.0276, 0.0504, -0.7305, 0.5015])
mm ax1 tensor([ 0.3166, -0.4357, -1.0276])
ms none tensor(-1.3849)
ms ax1 tensor([ 0.9499, -1.3072, -1.0276])
lse tensor([3.7453, 2.9333])
lsm tensor([[-5.3180, -3.8686, -0.1584, -5.5766, -2.1466],
        [-4.2103, -2.6078, -3.4124, -1.5543, -0.4047]])
tlp tensor([[-2.6539, -1.9474, -2.8140],
        [-1.6476, -1.2624, -3.1213]])
slp tensor([-7.4154, -2.9100])
seqlp F tensor([-7.4154, -2.9100])
seqlp T tensor([-2.4718, -1.4550])
entropy tensor(1.3433)
normalize tensor([-1.8910, -0.0302, 0.8794, 0.6768, 0.4327, -0.0676])

Probe normalize/entropy conventions

shell
cat > /app/_p2.py << 'EOF'
import torch, reforacle as R
x=torch.tensor([1.,2.,3.,4.])
print("normalize eps1e-8", R.normalize(x,1e-8))
mu=x.mean(); 
print("  (x-mu)/std(unbiased)", (x-mu)/x.std(unbiased=True))
print("  (x-mu)/std(biased)", (x-mu)/x.std(unbiased=False))
print("  (x-mu)/(std_b+eps)", (x-mu)/(x.std(unbiased=False)+1e-8))
# entropy check with simple logits
logits=torch.tensor([[[0.,0.],[10.,-10.]]])  # 1,2,2
mask=torch.tensor([[1.,1.]])
p=torch.softmax(logits,-1); ent=-(p*torch.log_softmax(logits,-1)).sum(-1)
print("per token ent", ent)
print("R.entropy full mask", R.entropy(logits,mask))
mask2=torch.tensor([[1.,0.]])
print("R.entropy half mask", R.entropy(logits,mask2))
# masked_mean empty axis
v=torch.tensor([[1.,2.],[3.,4.]]); m=torch.tensor([[0.,0.],[1.,1.]])
print("mm ax1 empty row", R.masked_mean(v,m,axis=1))
EOF
python3 /app/_p2.py
output
normalize eps1e-8 tensor([-1.1619, -0.3873, 0.3873, 1.1619])
  (x-mu)/std(unbiased) tensor([-1.1619, -0.3873, 0.3873, 1.1619])
  (x-mu)/std(biased) tensor([-1.3416, -0.4472, 0.4472, 1.3416])
  (x-mu)/(std_b+eps) tensor([-1.3416, -0.4472, 0.4472, 1.3416])
per token ent tensor([[6.9315e-01, 4.1223e-08]])
R.entropy full mask tensor(0.3466)
R.entropy half mask tensor(0.6931)
mm ax1 empty row tensor([0.0000, 3.5000])

[stdout]
normalize eps1e-8 tensor([-1.1619, -0.3873, 0.3873, 1.1619])
  (x-mu)/std(unbiased) tensor([-1.1619, -0.3873, 0.3873, 1.1619])
  (x-mu)/std(biased) tensor([-1.3416, -0.4472, 0.4472, 1.3416])
  (x-mu)/(std_b+eps) tensor([-1.3416, -0.4472, 0.4472, 1.3416])
per token ent tensor([[6.9315e-01, 4.1223e-08]])
R.entropy full mask tensor(0.3466)
R.entropy half mask tensor(0.6931)
mm ax1 empty row tensor([0.0000, 3.5000])

Probe KL and RL functions

shell
cat > /app/_p3.py << 'EOF'
import torch, reforacle as R
lp=torch.tensor([-1.0,-2.0,-0.5]); rlp=torch.tensor([-1.5,-1.0,-0.7])
print("reverse_kl", R.reverse_kl(lp,rlp))
print("symmetric_kl", R.symmetric_kl(lp,rlp))
for est in ["k1","k2","k3"]:
    try: print("kl_penalty",est, R.kl_penalty(lp,rlp,est))
    except Exception as e: print(est,"ERR",e)
print("importance_ratio noclip", R.importance_ratio(lp,rlp,None))
print("importance_ratio clip0.1", R.importance_ratio(lp,rlp,0.1))
r=torch.tensor([1.0,2.0,3.0])
print("discounted_returns g0.9", R.discounted_returns(r,0.9))
# gae
rewards=torch.tensor([1.0,0.0,-1.0]); values=torch.tensor([0.5,0.2,0.1])
print("gae", R.gae(rewards,values,torch.tensor(0.3),0.99,0.95))
print("lambda_returns", R.lambda_returns(rewards,values,torch.tensor(0.3),0.99,0.95))
v=torch.tensor([1.,2.,3.,4.,5.]); m=torch.tensor([1.,1.,0.,1.,1.])
print("whiten shiftT", R.whiten(v,m,True))
print("whiten shiftF", R.whiten(v,m,False))
print("masked_whiten shiftT", R.masked_whiten(v,m,True))
print("masked_whiten shiftF", R.masked_whiten(v,m,False))
EOF
python3 /app/_p3.py
output
reverse_kl tensor([0.1065, 0.7183, 0.0187])
symmetric_kl tensor([0.1276, 0.5431, 0.0201])
kl_penalty k1 tensor([ 0.5000, -1.0000, 0.2000])
kl_penalty k2 tensor([0.1250, 0.5000, 0.0200])
kl_penalty k3 tensor([0.1065, 0.7183, 0.0187])
importance_ratio noclip tensor([1.6487, 0.3679, 1.2214])
importance_ratio clip0.1 tensor([1.1000, 0.9000, 1.1000])
discounted_returns g0.9 tensor([5.2300, 4.7000, 3.0000])
gae tensor([-0.1073, -0.8562, -0.8030])
lambda_returns tensor([ 0.3927, -0.6562, -0.7030])
whiten shiftT tensor([-1.2649, -0.6325, 0.0000, 0.6325, 1.2649])
whiten shiftF tensor([1.7351, 2.3675, 3.0000, 3.6325, 4.2649])
masked_whiten shiftT tensor([-1.2649, -0.6325, 0.0000, 0.6325, 1.2649])
masked_whiten shiftF tensor([1.7351, 2.3675, 3.0000, 3.6325, 4.2649])

[stdout]
reverse_kl tensor([0.1065, 0.7183, 0.0187])
symmetric_kl tensor([0.1276, 0.5431, 0.0201])
kl_penalty k1 tensor([ 0.5000, -1.0000, 0.2000])
kl_penalty k2 tensor([0.1250, 0.5000, 0.0200])
kl_penalty k3 tensor([0.1065, 0.7183, 0.0187])
importance_ratio noclip tensor([1.6487, 0.3679, 1.2214])
importance_ratio clip0.1 tensor([1.1000, 0.9000, 1.1000])
discounted_returns g0.9 tensor([5.2300, 4.7000, 3.0000])
gae tensor([-0.1073, -0.8562, -0.8030])
lambda_returns tensor([ 0.3927, -0.6562, -0.7030])
whiten shiftT tensor([-1.2649, -0.6325, 0.0000, 0.6325, 1.2649])
whiten shiftF tensor([1.7351, 2.3675, 3.0000, 3.6325, 4.2649])
masked_whiten shiftT tensor([-1.2649, -0.6325, 0.0000, 0.6325, 1.2649])
masked_whiten shiftF tensor([1.7351, 2.3675, 3.0000, 3.6325, 4.2649])

Probe whiten/pg/value losses

shell
cat > /app/_p4.py << 'EOF'
import torch, reforacle as R
v=torch.tensor([1.,2.,3.,10.,20.]); m=torch.tensor([1.,1.,1.,0.,0.])
print("whiten T", R.whiten(v,m,True))
print("masked_whiten T", R.masked_whiten(v,m,True))
# manual masked
mv=v[m.bool()]; 
print("  masked mean/std", mv.mean(), mv.std(unbiased=True), mv.std(unbiased=False))
print("  full mean/std", v.mean(), v.std(unbiased=True))
# clipped_pg_loss
logp=torch.tensor([-1.0,-2.0,-0.5,-1.2]); old=torch.tensor([-1.1,-1.5,-0.5,-1.0])
adv=torch.tensor([1.0,-1.0,2.0,-0.5]); mask=torch.tensor([1.,1.,1.,0.])
print("clipped_pg", R.clipped_pg_loss(logp,old,adv,mask,0.2,0.2))
print("clip_fraction", R.clip_fraction(logp,old,0.2))
# value_loss
values=torch.tensor([1.,2.,3.]); oldv=torch.tensor([1.1,1.8,3.2]); ret=torch.tensor([1.5,2.5,2.0])
print("value_loss clip0.2", R.value_loss(values,oldv,ret,0.2))
print("value_loss clipNone", R.value_loss(values,oldv,ret,None))
print("huber", R.huber_value_loss(values,ret,1.0))
EOF
python3 /app/_p4.py
output
Exit code 1
whiten T tensor([-1.2247, 0.0000, 1.2247, 9.7980, 22.0454])
masked_whiten T tensor([-1.2247, 0.0000, 1.2247, 0.0000, 0.0000])
  masked mean/std tensor(2.) tensor(1.) tensor(0.8165)
  full mean/std tensor(7.2000) tensor(7.9812)
clipped_pg tensor(-0.7684)
clip_fraction tensor(0.2500)
value_loss clip0.2 tensor(0.2500)
Traceback (most recent call last):
  File "/app/_p4.py", line 17, in <module>
    print("value_loss clipNone", R.value_loss(values,oldv,ret,None))
                                 ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "reforacle.py", line 109, in reforacle.value_loss
TypeError: bad operand type for unary -: 'NoneType'

[error] tool reported failure

Probe huber and group advantages

shell
cat > /app/_p5.py << 'EOF'
import torch, reforacle as R
values=torch.tensor([1.,2.,3.]); ret=torch.tensor([1.5,2.5,5.0])
print("huber d1", R.huber_value_loss(values,ret,1.0))
# manual huber: 0.5*mean of huber
d=values-ret; a=d.abs(); delta=1.0
h=torch.where(a<=delta,0.5*d*d,delta*(a-0.5*delta))
print("  mean huber", h.mean(), "0.5*mean", 0.5*h.mean())
# group advantages
rewards=torch.tensor([1.,2.,3.,4.,5.,6.])
print("grpo scaleT", R.grpo_advantages(rewards,3,True))
print("grpo scaleF", R.grpo_advantages(rewards,3,False))
print("rloo", R.rloo_advantages(rewards,3))
print("group_mean_baseline", R.group_mean_baseline(rewards,3))
adv=torch.tensor([1.,2.,3.,4.]); m=torch.tensor([1.,1.,0.,1.])
print("advantage_mean_std", R.advantage_mean_std(adv,m))
EOF
python3 /app/_p5.py
output
huber d1 tensor(0.5833)
  mean huber tensor(0.5833) 0.5*mean tensor(0.2917)
grpo scaleT tensor([-1.2247, 0.0000, 1.2247, -1.2247, 0.0000, 1.2247])
grpo scaleF tensor([-1., 0., 1., -1., 0., 1.])
rloo tensor([-1.5000, 0.0000, 1.5000, -1.5000, 0.0000, 1.5000])
group_mean_baseline tensor([-1., 0., 1., -1., 0., 1.])
advantage_mean_std tensor([2.3333, 1.2472])

[stdout]
huber d1 tensor(0.5833)
  mean huber tensor(0.5833) 0.5*mean tensor(0.2917)
grpo scaleT tensor([-1.2247, 0.0000, 1.2247, -1.2247, 0.0000, 1.2247])
grpo scaleF tensor([-1., 0., 1., -1., 0., 1.])
rloo tensor([-1.5000, 0.0000, 1.5000, -1.5000, 0.0000, 1.5000])
group_mean_baseline tensor([-1., 0., 1., -1., 0., 1.])
advantage_mean_std tensor([2.3333, 1.2472])

Probe dpo/ce/entropy

shell
cat > /app/_p6.py << 'EOF'
import torch, reforacle as R
r=R.advantage_mean_std(torch.tensor([1.,2.,4.]),torch.tensor([1.,1.,1.]))
print("ams type", type(r), r)
# dpo/ipo
pc,pr,rc,rr=torch.tensor([-2.]),torch.tensor([-3.]),torch.tensor([-2.5]),torch.tensor([-2.2])
print("dpo ls0", R.dpo_loss(pc,pr,rc,rr,0.1,0.0))
print("dpo ls0.1", R.dpo_loss(pc,pr,rc,rr,0.1,0.1))
print("ipo", R.ipo_loss(pc,pr,rc,rr,0.1))
print("bt", R.bradley_terry_logit(torch.tensor([1.0]),torch.tensor([0.3]),0.5))
# smoothed_nll / cross_entropy
logits=torch.tensor([[2.,1.,0.,-1.],[0.,0.,1.,2.]]); labels=torch.tensor([0,3])
print("smoothed0", R.smoothed_nll(logits,labels,0.0))
print("smoothed0.1", R.smoothed_nll(logits,labels,0.1))
print("ce", R.cross_entropy(logits,labels,-100))
labels2=torch.tensor([0,-100])
print("ce ignore", R.cross_entropy(logits,labels2,-100))
# normalized_entropy
lg=torch.tensor([[[0.,0.,0.,0.]]]); mk=torch.tensor([[1.]])
print("norm_ent uniform", R.normalized_entropy(lg,mk))
print("logprob_at_temp", R.logprob_at_temperature(logits.unsqueeze(0),labels.unsqueeze(0),2.0))
EOF
python3 /app/_p6.py
output
ams type <class 'torch.Tensor'> tensor([2.3333, 1.2472])
dpo ls0 tensor(0.6303)
dpo ls0.1 tensor(0.6433)
ipo tensor(13.6900)
bt tensor([0.3500])
smoothed0 tensor(0.4670)
smoothed0.1 tensor(0.6045)
ce tensor(0.4670)
ce ignore tensor(0.4402)
norm_ent uniform tensor(1.)
logprob_at_temp tensor([[-0.7873, -0.8511]])

[stdout]
ams type <class 'torch.Tensor'> tensor([2.3333, 1.2472])
dpo ls0 tensor(0.6303)
dpo ls0.1 tensor(0.6433)
ipo tensor(13.6900)
bt tensor([0.3500])
smoothed0 tensor(0.4670)
smoothed0.1 tensor(0.6045)
ce tensor(0.4670)
ce ignore tensor(0.4402)
norm_ent uniform tensor(1.)
logprob_at_temp tensor([[-0.7873, -0.8511]])

Probe masks and simple ops

shell
cat > /app/_p7.py << 'EOF'
import torch, reforacle as R
probs=torch.tensor([[0.1,0.2,0.3,0.4],[0.5,0.3,0.15,0.05]])
print("top_p 0.7\n", R.top_p_mask(probs,0.7))
logits=torch.tensor([[2.,1.,0.,-1.],[0.,3.,1.,2.]])
print("top_k 2\n", R.top_k_mask(logits,2))
print("argmax", R.argmax_tokens(logits))
print("mode_label", R.mode_label(torch.tensor([1,2,2,3,3,3])))
print("median odd", R.median_reward(torch.tensor([3.,1.,2.,5.,4.])))
print("median even", R.median_reward(torch.tensor([1.,2.,3.,4.])))
v=torch.tensor([1.,5.,2.,8.]); m=torch.tensor([1.,0.,1.,1.])
print("masked_max", R.masked_max(v,m))
print("quantile_lower 0.25", R.quantile_lower(torch.tensor([1.,2.,3.,4.,5.]),0.25))
print("pad_mask", R.pad_mask_from_lengths(torch.tensor([2,3,1]),4))
print("first_nonzero", R.first_nonzero_index(torch.tensor([[0.,0.,1.,0.],[1.,0.,0.,0.],[0.,0.,0.,0.]])))
print("cummax", R.cumulative_max(torch.tensor([1.,3.,2.,5.,4.])))
print("bucketize", R.bucketize_reward(torch.tensor([0.5,1.5,2.5,3.5]),torch.tensor([1.,2.,3.])))
EOF
python3 /app/_p7.py
output
top_p 0.7
 tensor([[0.0000, 0.0000, 0.4286, 0.5714],
        [0.6250, 0.3750, 0.0000, 0.0000]])
top_k 2
 tensor([[ True, True, False, False],
        [False, True, False, True]])
argmax tensor([0, 1])
mode_label tensor(3)
median odd tensor(3.)
median even tensor(2.5000)
masked_max tensor(8.)
quantile_lower 0.25 tensor(2.)
pad_mask tensor([[1., 1., 0., 0.],
        [1., 1., 1., 0.],
        [1., 0., 0., 0.]])
first_nonzero tensor([ 2, 0, -1])
cummax tensor([1., 1., 3., 3., 5.])
bucketize tensor([0, 1, 2, 3])

[stdout]
top_p 0.7
 tensor([[0.0000, 0.0000, 0.4286, 0.5714],
        [0.6250, 0.3750, 0.0000, 0.0000]])
top_k 2
 tensor([[ True, True, False, False],
        [False, True, False, True]])
argmax tensor([0, 1])
mode_label tensor(3)
median odd tensor(3.)
median even tensor(2.5000)
masked_max tensor(8.)
quantile_lower 0.25 tensor(2.)
pad_mask tensor([[1., 1., 0., 0.],
        [1., 1., 1., 0.],
        [1., 0., 0., 0.]])
first_nonzero tensor([ 2, 0, -1])
cummax tensor([1., 1., 3., 3., 5.])
bucketize tensor([0, 1, 2, 3])

Probe edge cases

shell
cat > /app/_p8.py << 'EOF'
import torch, reforacle as R
print("cummax A", R.cumulative_max(torch.tensor([1.,3.,2.,5.,4.])))
print("cummax B", R.cumulative_max(torch.tensor([5.,1.,2.,3.])))
print("cummax 2d", R.cumulative_max(torch.tensor([[1.,3.,2.],[4.,0.,9.]])))
print("quantile 0.3 on 1..5", R.quantile_lower(torch.tensor([1.,2.,3.,4.,5.]),0.3))
print("quantile 0.5 on 1..4", R.quantile_lower(torch.tensor([1.,2.,3.,4.]),0.5))
print("bucketize eq", R.bucketize_reward(torch.tensor([1.,2.,3.]),torch.tensor([1.,2.,3.])))
print("mode tie", R.mode_label(torch.tensor([1,1,2,2])))
# top_p single dominant
print("top_p dom", R.top_p_mask(torch.tensor([[0.8,0.1,0.1]]),0.5))
print("top_p p1", R.top_p_mask(torch.tensor([[0.4,0.3,0.2,0.1]]),1.0))
EOF
python3 /app/_p8.py
output
cummax A tensor([1., 1., 3., 3., 5.])
cummax B tensor([5., 5., 5., 5.])
cummax 2d tensor([[1., 3., 2.],
        [1., 3., 2.]])
quantile 0.3 on 1..5 tensor(2.)
quantile 0.5 on 1..4 tensor(2.)
bucketize eq tensor([1, 2, 3])
mode tie tensor(2)
top_p dom tensor([[1., 0., 0.]])
top_p p1 tensor([[0.4000, 0.3000, 0.2000, 0.1000]])

[stdout]
cummax A tensor([1., 1., 3., 3., 5.])
cummax B tensor([5., 5., 5., 5.])
cummax 2d tensor([[1., 3., 2.],
        [1., 3., 2.]])
quantile 0.3 on 1..5 tensor(2.)
quantile 0.5 on 1..4 tensor(2.)
bucketize eq tensor([1, 2, 3])
mode tie tensor(2)
top_p dom tensor([[1., 0., 0.]])
top_p p1 tensor([[0.4000, 0.3000, 0.2000, 0.1000]])

Probe normalized_entropy, reverse_kl_objective

shell
cat > /app/_p9.py << 'EOF'
import torch, reforacle as R
g=torch.Generator().manual_seed(5)
def rn(*s): return torch.randn(*s,generator=g)
# normalized_entropy convention
lg=rn(1,3,5); mk=torch.tensor([[1.,1.,0.]])
print("norm_ent", R.normalized_entropy(lg,mk))
print("entropy/log5", R.entropy(lg,mk)/torch.log(torch.tensor(5.0)))
# importance_ratio clip None
print("ir None", R.importance_ratio(torch.tensor([0.5]),torch.tensor([0.0]),None))
# reverse_kl_objective
B,T,V=2,3,4
logits=rn(B,T,V); ref=rn(B,T,V); labels=torch.randint(0,V,(B,T),generator=g)
mask=(torch.rand(B,T,generator=g)>0.3).float(); adv=rn(B,T)
print("rkl_obj", R.reverse_kl_objective(logits,ref,labels,mask,adv,0.1))
EOF
python3 /app/_p9.py
output
norm_ent tensor(0.7432)
entropy/log5 tensor(0.7432)
ir None tensor([1.6487])
rkl_obj tensor(2.1034)

[stdout]
norm_ent tensor(0.7432)
entropy/log5 tensor(0.7432)
ir None tensor([1.6487])
rkl_obj tensor(2.1034)

Solve reverse_kl_objective

shell
cat > /app/_p10.py << 'EOF'
import torch, reforacle as R
g=torch.Generator().manual_seed(5)
def rn(*s): return torch.randn(*s,generator=g)
B,T,V=2,3,4
logits=rn(B,T,V); ref=rn(B,T,V); labels=torch.randint(0,V,(B,T),generator=g)
mask=(torch.rand(B,T,generator=g)>0.3).float(); adv=rn(B,T)

def token_logprobs(lg,lb):
    lp=torch.log_softmax(lg,-1)
    return lp.gather(-1,lb.unsqueeze(-1)).squeeze(-1)
def mmean(v,m):
    return (v*m).sum()/m.sum().clamp(min=1e-9) if m.sum()>0 else torch.tensor(0.)

logp=token_logprobs(logits,labels)
reflp=token_logprobs(ref,labels)
d=logp-reflp
kl=torch.exp(-d)+d-1  # reverse_kl per token
target=R.reverse_kl_objective(logits,ref,labels,mask,adv,0.1)
print("target",target)
for name,cand in [
 ("-adv*logp+beta*kl", (-adv*logp+0.1*kl)),
 ("-(adv*logp)-beta*kl", (-adv*logp-0.1*kl)),
 ("adv*logp+beta*kl", (adv*logp+0.1*kl)),
]:
    print(name, mmean(cand,mask).item())
EOF
python3 /app/_p10.py
output
target tensor(-0.5832)
-adv*logp+beta*kl -0.24233147501945496
-(adv*logp)-beta*kl -0.840238094329834
adv*logp+beta*kl 0.840238094329834

[stdout]
target tensor(-0.5832)
-adv*logp+beta*kl -0.24233147501945496
-(adv*logp)-beta*kl -0.840238094329834
adv*logp+beta*kl 0.840238094329834

Break down components

shell
cat >> /app/_p10.py << 'EOF'

print("mm(-adv*logp)", mmean(-adv*logp,mask).item())
print("mm(adv*logp)", mmean(adv*logp,mask).item())
print("mm(kl)", mmean(kl,mask).item())
print("mm(logp)", mmean(logp,mask).item())
# maybe -adv*logp - beta*... or +? target -0.5832
# try beta*kl - adv*logp with various
import itertools
base_adv=mmean(-adv*logp,mask).item()
base_kl=mmean(kl,mask).item()
print("check a*base_adv? target/base", -0.5832)
print("mm(-adv*logp)+0.1*mm(kl)", base_adv+0.1*base_kl)
print("mm(-adv*logp)-0.1*mm(kl)", base_adv-0.1*base_kl)
EOF
python3 /app/_p10.py
output
target tensor(-0.5832)
-adv*logp+beta*kl -0.24233147501945496
-(adv*logp)-beta*kl -0.840238094329834
adv*logp+beta*kl 0.840238094329834
mm(-adv*logp) -0.5412847995758057
mm(adv*logp) 0.5412847995758057
mm(kl) 2.989532947540283
mm(logp) -2.2856521606445312
check a*base_adv? target/base -0.5832
mm(-adv*logp)+0.1*mm(kl) -0.24233150482177734
mm(-adv*logp)-0.1*mm(kl) -0.840238094329834

[stdout]
target tensor(-0.5832)
-adv*logp+beta*kl -0.24233147501945496
-(adv*logp)-beta*kl -0.840238094329834
adv*logp+beta*kl 0.840238094329834
mm(-adv*logp) -0.5412847995758057
mm(adv*logp) 0.5412847995758057
mm(kl) 2.989532947540283
mm(logp) -2.2856521606445312
check a*base_adv? target/base -0.5832
mm(-adv*logp)+0.1*mm(kl) -0.24233150482177734
mm(-adv*logp)-0.1*mm(kl) -0.840238094329834

More KL candidates

shell
cat >> /app/_p10.py << 'EOF'

fwd=torch.exp(d)-d-1
k1=d; k2=0.5*d*d
print("mm(k1)", mmean(k1,mask).item())
print("mm(fwd)", mmean(fwd,mask).item())
print("mm(k2)", mmean(k2,mask).item())
print("need per-tok mm ~", (-0.5832-(-0.5412847995758057))/0.1)
# maybe adv includes ref: pg with ratio to ref?
ratio=torch.exp(logp-reflp)
print("mm(-adv*ratio)", mmean(-adv*ratio,mask).item())
print("mm(-adv*ratio+0.1kl)", mmean(-adv*ratio+0.1*kl,mask).item())
EOF
python3 /app/_p10.py 2>&1 | tail -6
output
mm(k1) -0.8586515188217163
mm(fwd) 0.9142506122589111
mm(k2) 1.3712430000305176
need per-tok mm ~ -0.4191520042419439
mm(-adv*ratio) 0.45195621252059937
mm(-adv*ratio+0.1kl) 0.7509095668792725

[stdout]
mm(k1) -0.8586515188217163
mm(fwd) 0.9142506122589111
mm(k2) 1.3712430000305176
need per-tok mm ~ -0.4191520042419439
mm(-adv*ratio) 0.45195621252059937
mm(-adv*ratio+0.1kl) 0.7509095668792725

Log-ratio candidates

shell
cat >> /app/_p10.py << 'EOF'

print("mm(-adv*d)", mmean(-adv*d,mask).item())
print("mm(-adv*reflp)", mmean(-adv*reflp,mask).item())
print("mm(-adv*d)+0.1*mm(kl)", mmean(-adv*d,mask).item()+0.1*mmean(kl,mask).item())
print("mm(-adv*d + 0.1*kl)", mmean(-adv*d+0.1*kl,mask).item())
print("mm(-adv*d - 0.1*kl)", mmean(-adv*d-0.1*kl,mask).item())
print("mm(adv*d + 0.1*kl)", mmean(adv*d+0.1*kl,mask).item())
# ratio pg with kl subtracted in advantage
print("mm(-(adv-0.1*kl)*logp)", mmean(-(adv-0.1*kl)*logp,mask).item())
EOF
python3 /app/_p10.py 2>&1 | tail -6
output
mm(-adv*reflp) -0.4299217164516449
mm(-adv*d)+0.1*mm(kl) 0.18759019672870636
mm(-adv*d + 0.1*kl) 0.18759021162986755
mm(-adv*d - 0.1*kl) -0.4103164076805115
mm(adv*d + 0.1*kl) 0.4103164076805115
mm(-(adv-0.1*kl)*logp) -1.5442219972610474

[stdout]
mm(-adv*reflp) -0.4299217164516449
mm(-adv*d)+0.1*mm(kl) 0.18759019672870636
mm(-adv*d + 0.1*kl) 0.18759021162986755
mm(-adv*d - 0.1*kl) -0.4103164076805115
mm(adv*d + 0.1*kl) 0.4103164076805115
mm(-(adv-0.1*kl)*logp) -1.5442219972610474

Isolate KL and adv terms

shell
cat > /app/_p11.py << 'EOF'
import torch, reforacle as R
g=torch.Generator().manual_seed(5)
def rn(*s): return torch.randn(*s,generator=g)
B,T,V=2,3,4
logits=rn(B,T,V); ref=rn(B,T,V); labels=torch.randint(0,V,(B,T),generator=g)
mask=torch.ones(B,T); adv=rn(B,T)
def tlp(lg,lb):
    return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
logp=tlp(logits,labels); reflp=tlp(ref,labels); d=logp-reflp
kl=torch.exp(-d)+d-1
def mm(v,m): return (v*m).sum()/m.sum()
print("adv=0:", R.reverse_kl_objective(logits,ref,labels,mask,torch.zeros(B,T),0.1).item())
print(" 0.1*mm(kl)=",0.1*mm(kl,mask).item())
print(" 0.1*mm(k1)=",0.1*mm(d,mask).item())
print("beta=0:", R.reverse_kl_objective(logits,ref,labels,mask,adv,0.0).item())
print(" mm(-adv*logp)=",mm(-adv*logp,mask).item())
print(" mm(-adv*d)=",mm(-adv*d,mask).item())
print(" mm(adv*logp)=",mm(adv*logp,mask).item())
EOF
python3 /app/_p11.py
output
adv=0: 0.7638916969299316
 0.1*mm(kl)= 0.3819458961486817
 0.1*mm(k1)= -0.08011066913604736
beta=0: 0.21610014140605927
 mm(-adv*logp)= -0.39875224232673645
 mm(-adv*d)= -0.7132778763771057
 mm(adv*logp)= 0.39875224232673645

[stdout]
adv=0: 0.7638916969299316
 0.1*mm(kl)= 0.3819458961486817
 0.1*mm(k1)= -0.08011066913604736
beta=0: 0.21610014140605927
 mm(-adv*logp)= -0.39875224232673645
 mm(-adv*d)= -0.7132778763771057
 mm(adv*logp)= 0.39875224232673645

True KL and ratio

shell
cat >> /app/_p11.py << 'EOF'

lpp=torch.log_softmax(logits,-1); lpr=torch.log_softmax(ref,-1)
pp=lpp.exp()
trueKL=(pp*(lpp-lpr)).sum(-1)  # reverse KL(policy||ref) per token
print("--- true distributional ---")
print("0.1*mm(trueKL)=",0.1*mm(trueKL,mask).item(), " x2=",0.2*mm(trueKL,mask).item())
ratio=torch.exp(d)
print("mm(-adv*ratio)=",mm(-adv*ratio,mask).item())
print("mm(adv*ratio)=",mm(adv*ratio,mask).item())
print("mm(-adv*logp)+? target beta0=0.2161")
# maybe advantage term is -mm(adv*logp) but adv whitened? or masked_mean differ
EOF
python3 /app/_p11.py 2>&1 | tail -6
output
mm(adv*logp)= 0.39875224232673645
--- true distributional ---
0.1*mm(trueKL)= 0.08726447224617005  x2= 0.1745289444923401
mm(-adv*ratio)= -0.8311963081359863
mm(adv*ratio)= 0.8311963081359863
mm(-adv*logp)+? target beta0=0.2161

[stdout]
 mm(adv*logp)= 0.39875224232673645
--- true distributional ---
0.1*mm(trueKL)= 0.08726447224617005  x2= 0.1745289444923401
mm(-adv*ratio)= -0.8311963081359863
mm(adv*ratio)= 0.8311963081359863
mm(-adv*logp)+? target beta0=0.2161

Exact decomposition

shell
cat >> /app/_p11.py << 'EOF'

print("=== exact ===")
o0=R.reverse_kl_objective(logits,ref,labels,mask,torch.zeros(B,T),0.1).item()
print("obj adv0 beta0.1", o0, " /0.1=", o0/0.1)
print("sum(kl)",kl.sum().item(),"mean",kl.mean().item())
print("sum(trueKL)",trueKL.sum().item(),"mean",trueKL.mean().item())
# beta=0 case decompose
ob=R.reverse_kl_objective(logits,ref,labels,mask,adv,0.0).item()
print("obj beta0", ob)
print("mm(-adv*logp)",mm(-adv*logp,mask).item())
print("mm(-adv*reflp)",mm(-adv*reflp,mask).item())
# maybe adv gets whitened inside
aw=(adv-adv.mean())/adv.std()
print("mm(-aw*logp)",mm(-aw*logp,mask).item())
EOF
python3 /app/_p11.py 2>&1 | tail -9
output
mm(-adv*logp)+? target beta0=0.2161
=== exact ===
obj adv0 beta0.1 0.7638916969299316  /0.1= 7.638916969299316
sum(kl) 22.9167537689209 mean 3.8194589614868164
sum(trueKL) 5.235868453979492 mean 0.8726447224617004
obj beta0 0.21610014140605927
mm(-adv*logp) -0.39875224232673645
mm(-adv*reflp) 0.31452563405036926
mm(-aw*logp) -0.6378178000450134

[stdout]
mm(-adv*logp)+? target beta0=0.2161
=== exact ===
obj adv0 beta0.1 0.7638916969299316  /0.1= 7.638916969299316
sum(kl) 22.9167537689209 mean 3.8194589614868164
sum(trueKL) 5.235868453979492 mean 0.8726447224617004
obj beta0 0.21610014140605927
mm(-adv*logp) -0.39875224232673645
mm(-adv*reflp) 0.31452563405036926
mm(-aw*logp) -0.6378178000450134

Estimator means

shell
cat >> /app/_p11.py << 'EOF'

print("=== estimator means (mask all ones) ===")
for nm,val in [("k1",d),("k2",0.5*d*d),("k3rev",kl),("fwd",torch.exp(d)-d-1),
   ("sym",torch.cosh(d)-1),("abs_d",d.abs()),("d2",d*d)]:
    print(nm, mm(val,mask).item())
print("need 7.639")
EOF
python3 /app/_p11.py 2>&1 | tail -9
output
=== estimator means (mask all ones) ===
k1 -0.8011066913604736
k2 1.6306487321853638
k3rev 3.8194589614868164
fwd 1.1105259656906128
sym 2.4649922847747803
abs_d 1.612099289894104
d2 3.2612974643707275
need 7.639

[stdout]
=== estimator means (mask all ones) ===
k1 -0.8011066913604736
k2 1.6306487321853638
k3rev 3.8194589614868164
fwd 1.1105259656906128
sym 2.4649922847747803
abs_d 1.612099289894104
d2 3.2612974643707275
need 7.639

Reduction hypothesis test

shell
cat >> /app/_p11.py << 'EOF'

print("=== reduction test ===")
def mm_ax1(v,m): return (v*m).sum(1)/m.sum(1)
kterm=mm_ax1(kl,mask)
print("per-seq kl mean", kterm, "sum", kterm.sum().item())
aterm=mm_ax1(-adv*logp,mask)
print("per-seq adv mean", aterm, "sum", aterm.sum().item(), "mean", aterm.mean().item())
print("obj beta0 target 0.21610")
# combined full
full=R.reverse_kl_objective(logits,ref,labels,mask,adv,0.1).item()
print("full target", full)
print("cand sum: ", (mm_ax1(-adv*logp,mask)+0.1*mm_ax1(kl,mask)).sum().item())
print("cand mean:", (mm_ax1(-adv*logp,mask)+0.1*mm_ax1(kl,mask)).mean().item())
EOF
python3 /app/_p11.py 2>&1 | tail -7
output
=== reduction test ===
per-seq kl mean tensor([7.2879, 0.3511]) sum 7.638917446136475
per-seq adv mean tensor([-1.4003, 0.6028]) sum -0.7975046038627625 mean -0.3987523019313812
obj beta0 target 0.21610
full target 0.9799919128417969
cand sum:  -0.033612847328186035
cand mean: -0.016806423664093018

[stdout]
=== reduction test ===
per-seq kl mean tensor([7.2879, 0.3511]) sum 7.638917446136475
per-seq adv mean tensor([-1.4003, 0.6028]) sum -0.7975046038627625 mean -0.3987523019313812
obj beta0 target 0.21610
full target 0.9799919128417969
cand sum:  -0.033612847328186035
cand mean: -0.016806423664093018

Brute force beta0 term

shell
cat >> /app/_p11.py << 'EOF'

print("=== beta0 brute (target 0.21610) ===")
ratio=torch.exp(d)
cands={
 "mean(-adv*logp)":(-adv*logp),
 "mean(adv*logp)":(adv*logp),
 "mean(-adv*d)":(-adv*d),
 "mean(-adv*ratio)":(-adv*ratio),
 "mean(adv*ratio)":(adv*ratio),
 "mean(-adv*reflp)":(-adv*reflp),
 "mean(-(adv)*(logp-reflp.detach))":(-adv*d),
}
for n,v in cands.items(): print(n, mm(v,mask).item(), " sum_ax1", mm_ax1(v,mask).sum().item())
EOF
python3 /app/_p11.py 2>&1 | tail -8
output
=== beta0 brute (target 0.21610) ===
mean(-adv*logp) -0.39875224232673645  sum_ax1 -0.7975046038627625
mean(adv*logp) 0.39875224232673645  sum_ax1 0.7975046038627625
mean(-adv*d) -0.7132778763771057  sum_ax1 -1.426555871963501
mean(-adv*ratio) -0.8311963081359863  sum_ax1 -1.6623926162719727
mean(adv*ratio) 0.8311963081359863  sum_ax1 1.6623926162719727
mean(-adv*reflp) 0.31452563405036926  sum_ax1 0.6290512681007385
mean(-(adv)*(logp-reflp.detach)) -0.7132778763771057  sum_ax1 -1.426555871963501

[stdout]
=== beta0 brute (target 0.21610) ===
mean(-adv*logp) -0.39875224232673645  sum_ax1 -0.7975046038627625
mean(adv*logp) 0.39875224232673645  sum_ax1 0.7975046038627625
mean(-adv*d) -0.7132778763771057  sum_ax1 -1.426555871963501
mean(-adv*ratio) -0.8311963081359863  sum_ax1 -1.6623926162719727
mean(adv*ratio) 0.8311963081359863  sum_ax1 1.6623926162719727
mean(-adv*reflp) 0.31452563405036926  sum_ax1 0.6290512681007385
mean(-(adv)*(logp-reflp.detach)) -0.7132778763771057  sum_ax1 -1.426555871963501

Whitened advantage term

shell
cat >> /app/_p11.py << 'EOF'

print("=== whitened adv (target A=0.21610) ===")
am=adv.mean(); asb=adv.std(unbiased=False); asu=adv.std(unbiased=True)
for nm,aa in [("biased",(adv-am)/asb),("unbiased",(adv-am)/asu),("noshift_b",adv/asb)]:
  print(nm,"mean(-a*logp)",mm(-aa*logp,mask).item()," mean(a*logp)",mm(aa*logp,mask).item(),
        " sumax1(-a*logp)",mm_ax1(-aa*logp,mask).sum().item())
EOF
python3 /app/_p11.py 2>&1 | tail -4
output
=== whitened adv (target A=0.21610) ===
biased mean(-a*logp) -0.6986944079399109  mean(a*logp) 0.6986944079399109  sumax1(-a*logp) -1.3973889350891113
unbiased mean(-a*logp) -0.6378178000450134  mean(a*logp) 0.6378178000450134  sumax1(-a*logp) -1.2756356000900269
noshift_b mean(-a*logp) -0.554550051689148  mean(a*logp) 0.554550051689148  sumax1(-a*logp) -1.109100103378296

[stdout]
=== whitened adv (target A=0.21610) ===
biased mean(-a*logp) -0.6986944079399109  mean(a*logp) 0.6986944079399109  sumax1(-a*logp) -1.3973889350891113
unbiased mean(-a*logp) -0.6378178000450134  mean(a*logp) 0.6378178000450134  sumax1(-a*logp) -1.2756356000900269
noshift_b mean(-a*logp) -0.554550051689148  mean(a*logp) 0.554550051689148  sumax1(-a*logp) -1.109100103378296

Per-sequence A

shell
cat >> /app/_p11.py << 'EOF'

print("=== per-seq A decomposition ===")
adv0=adv.clone(); adv0[1]=0
adv1=adv.clone(); adv1[0]=0
A0=R.reverse_kl_objective(logits,ref,labels,mask,adv0,0.0).item()
A1=R.reverse_kl_objective(logits,ref,labels,mask,adv1,0.0).item()
print("A0",A0,"A1",A1,"sum",A0+A1,"target A",0.21610014)
print("row0 -adv*logp mean", (-adv[0]*logp[0]).mean().item(), " row1", (-adv[1]*logp[1]).mean().item())
print("row0 adv", adv[0].tolist(), "logp0", logp[0].tolist())
EOF
python3 /app/_p11.py 2>&1 | tail -5
output
noshift_b mean(-a*logp) -0.554550051689148  mean(a*logp) 0.554550051689148  sumax1(-a*logp) -1.109100103378296
=== per-seq A decomposition ===
A0 -0.9549288749694824 A1 1.1710290908813477 sum 0.21610021591186523 target A 0.21610014
row0 -adv*logp mean -1.4003394842147827  row1 0.6028348803520203
row0 adv [0.31667575240135193, -1.1957919597625732, -0.31266358494758606] logp0 [-3.8347020149230957, -3.6795475482940674, -3.2475860118865967]

[stdout]
noshift_b mean(-a*logp) -0.554550051689148  mean(a*logp) 0.554550051689148  sumax1(-a*logp) -1.109100103378296
=== per-seq A decomposition ===
A0 -0.9549288749694824 A1 1.1710290908813477 sum 0.21610021591186523 target A 0.21610014
row0 -adv*logp mean -1.4003394842147827  row1 0.6028348803520203
row0 adv [0.31667575240135193, -1.1957919597625732, -0.31266358494758606] logp0 [-3.8347020149230957, -3.6795475482940674, -3.2475860118865967]

Row0 detailed candidates

shell
cat >> /app/_p11.py << 'EOF'

print("=== row0 detail, A0=-0.95493 ===")
r0=torch.exp(d[0]); 
print("reflp0", reflp[0].tolist(), "ratio0", r0.tolist(), "d0", d[0].tolist())
a0=adv[0]
for nm,v in [("-a*logp",-a0*logp[0]),("-a*ratio",-a0*r0),("a*ratio",a0*r0),
   ("-a*d",-a0*d[0]),("a*d",a0*d[0]),("-a*reflp",-a0*reflp[0])]:
    print(nm,"mean",v.mean().item(),"sum",v.sum().item())
EOF
python3 /app/_p11.py 2>&1 | tail -7
output
reflp0 [-1.1920806169509888, -1.5303988456726074, -1.0226691961288452] ratio0 [0.07117444276809692, 0.11658336222171783, 0.10807639360427856] d0 [-2.6426215171813965, -2.14914870262146, -2.224916934967041]
-a*logp mean -1.4003394842147827 sum -4.201018333435059
-a*ratio mean 0.05022059381008148 sum 0.15066178143024445
a*ratio mean -0.05022059381008148 sum -0.15066178143024445
-a*d mean -0.8095770478248596 sum -2.4287312030792236
a*d mean 0.8095770478248596 sum 2.4287312030792236
-a*reflp mean -0.5907623171806335 sum -1.7722870111465454

[stdout]
reflp0 [-1.1920806169509888, -1.5303988456726074, -1.0226691961288452] ratio0 [0.07117444276809692, 0.11658336222171783, 0.10807639360427856] d0 [-2.6426215171813965, -2.14914870262146, -2.224916934967041]
-a*logp mean -1.4003394842147827 sum -4.201018333435059
-a*ratio mean 0.05022059381008148 sum 0.15066178143024445
a*ratio mean -0.05022059381008148 sum -0.15066178143024445
-a*d mean -0.8095770478248596 sum -2.4287312030792236
a*d mean 0.8095770478248596 sum 2.4287312030792236
-a*reflp mean -0.5907623171806335 sum -1.7722870111465454

Linearity and ref dependence

shell
cat >> /app/_p11.py << 'EOF'

print("=== structure probes ===")
print("A(adv)", R.reverse_kl_objective(logits,ref,labels,mask,adv,0.0).item())
print("A(2*adv)", R.reverse_kl_objective(logits,ref,labels,mask,2*adv,0.0).item())
print("A(adv), ref=logits", R.reverse_kl_objective(logits,logits,labels,mask,adv,0.0).item())
print("A(adv,beta.1) ref=logits", R.reverse_kl_objective(logits,logits,labels,mask,adv,0.1).item())
EOF
python3 /app/_p11.py 2>&1 | tail -5
output
=== structure probes ===
A(adv) 0.21610014140605927
A(2*adv) 0.43220028281211853
A(adv), ref=logits 0.21610014140605927
A(adv,beta.1) ref=logits 0.21610014140605927

[stdout]
=== structure probes ===
A(adv) 0.21610014140605927
A(2*adv) 0.43220028281211853
A(adv), ref=logits 0.21610014140605927
A(adv,beta.1) ref=logits 0.21610014140605927

Recover linear functional weights

shell
cat >> /app/_p11.py << 'EOF'

print("=== recover linear weights w ===")
W=torch.zeros(B,T)
for i in range(B):
  for t in range(T):
    e=torch.zeros(B,T); e[i,t]=1.0
    W[i,t]=R.reverse_kl_objective(logits,ref,labels,mask,e,0.0).item()
print("W\n",W)
print("logp\n",logp)
print("W/logp\n", W/logp)
print("-logp/6\n", -logp/6)
EOF
python3 /app/_p11.py 2>&1 | tail -12
output
W
 tensor([[0.8194, 0.8247, 0.7302],
        [0.8194, 0.8247, 0.7302]])
logp
 tensor([[-3.8347, -3.6795, -3.2476],
        [-1.0818, -1.2684, -1.1337]])
W/logp
 tensor([[-0.2137, -0.2241, -0.2248],
        [-0.7574, -0.6502, -0.6441]])
-logp/6
 tensor([[0.6391, 0.6133, 0.5413],
        [0.1803, 0.2114, 0.1889]])

[stdout]
W
 tensor([[0.8194, 0.8247, 0.7302],
        [0.8194, 0.8247, 0.7302]])
logp
 tensor([[-3.8347, -3.6795, -3.2476],
        [-1.0818, -1.2684, -1.1337]])
W/logp
 tensor([[-0.2137, -0.2241, -0.2248],
        [-0.7574, -0.6502, -0.6441]])
-logp/6
 tensor([[0.6391, 0.6133, 0.5413],
        [0.1803, 0.2114, 0.1889]])

Test W logit dependence

shell
cat >> /app/_p11.py << 'EOF'

print("=== does W depend on logits? ===")
g2=torch.Generator().manual_seed(99)
logits2=torch.randn(B,T,V,generator=g2)*3
W2=torch.zeros(B,T)
for i in range(B):
  for t in range(T):
    e=torch.zeros(B,T); e[i,t]=1.0
    W2[i,t]=R.reverse_kl_objective(logits2,ref,labels,mask,e,0.0).item()
print("W2\n",W2)
logp2=torch.log_softmax(logits2,-1).gather(-1,labels.unsqueeze(-1)).squeeze(-1)
print("logp2\n",logp2)
print("exp(logp2)\n", logp2.exp())
EOF
python3 /app/_p11.py 2>&1 | tail -8
output
tensor([[1.9100, 2.4247, 0.5501],
        [1.9100, 2.4247, 0.5501]])
logp2
 tensor([[ -0.9917, -8.7101, -0.1141],
        [-10.4685, -5.8379, -3.1863]])
exp(logp2)
 tensor([[3.7096e-01, 1.6490e-04, 8.9214e-01],
        [2.8418e-05, 2.9148e-03, 4.1325e-02]])

[stdout]
 tensor([[1.9100, 2.4247, 0.5501],
        [1.9100, 2.4247, 0.5501]])
logp2
 tensor([[ -0.9917, -8.7101, -0.1141],
        [-10.4685, -5.8379, -3.1863]])
exp(logp2)
 tensor([[3.7096e-01, 1.6490e-04, 8.9214e-01],
        [2.8418e-05, 2.9148e-03, 4.1325e-02]])

Pin reverse_kl_objective formula B=3

shell
cat > /app/_p12.py << 'EOF'
import torch, reforacle as R
g=torch.Generator().manual_seed(7)
B,T,V=3,2,4
logits=torch.randn(B,T,V,generator=g); ref=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g); mask=torch.ones(B,T)
logp=torch.log_softmax(logits,-1).gather(-1,labels.unsqueeze(-1)).squeeze(-1)
W=torch.zeros(B,T)
for i in range(B):
  for t in range(T):
    e=torch.zeros(B,T); e[i,t]=1.0
    W[i,t]=R.reverse_kl_objective(logits,ref,labels,mask,e,0.0).item()
print("W\n",W)
N=B*T
print("colsum logp /N: col-wise -(sum_i logp)/N")
print(-(logp.sum(0,keepdim=True)/N).expand(B,T))
print("-logp/N\n", -logp/N)
print("-logp/B\n", -logp/B)
EOF
python3 /app/_p12.py
output
W
 tensor([[0.7046, 0.5989],
        [0.7046, 0.5989],
        [0.7046, 0.5989]])
colsum logp /N: col-wise -(sum_i logp)/N
tensor([[0.7046, 0.5989],
        [0.7046, 0.5989],
        [0.7046, 0.5989]])
-logp/N
 tensor([[0.4869, 0.1417],
        [0.0587, 0.1134],
        [0.1590, 0.3437]])
-logp/B
 tensor([[0.9738, 0.2834],
        [0.1173, 0.2269],
        [0.3180, 0.6875]])

[stdout]
W
 tensor([[0.7046, 0.5989],
        [0.7046, 0.5989],
        [0.7046, 0.5989]])
colsum logp /N: col-wise -(sum_i logp)/N
tensor([[0.7046, 0.5989],
        [0.7046, 0.5989],
        [0.7046, 0.5989]])
-logp/N
 tensor([[0.4869, 0.1417],
        [0.0587, 0.1134],
        [0.1590, 0.3437]])
-logp/B
 tensor([[0.9738, 0.2834],
        [0.1173, 0.2269],
        [0.3180, 0.6875]])

Mask handling for reverse_kl_objective

shell
cat > /app/_p13.py << 'EOF'
import torch, reforacle as R
g=torch.Generator().manual_seed(7)
B,T,V=3,2,4
logits=torch.randn(B,T,V,generator=g); ref=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g)
mask=torch.tensor([[1.,1.],[1.,0.],[0.,1.]])
logp=torch.log_softmax(logits,-1).gather(-1,labels.unsqueeze(-1)).squeeze(-1)
W=torch.zeros(B,T)
for i in range(B):
  for t in range(T):
    e=torch.zeros(B,T); e[i,t]=1.0
    W[i,t]=R.reverse_kl_objective(logits,ref,labels,mask,e,0.0).item()
print("W (masked)\n",W)
Nsum=mask.sum().item()
# hypo: W_it = -(sum_j logp_jt*m_jt)/Nsum * m_it
colsum=(logp*mask).sum(0,keepdim=True)
print("cand -(colsum_logp)/sum(mask) * m\n", (-(colsum)/Nsum)*mask )
print("cand with per-col denom colmask")
colmask=mask.sum(0,keepdim=True)
print((-(colsum)/colmask.clamp(min=1))*mask)
EOF
python3 /app/_p13.py
output
W (masked)
 tensor([[0.8183, 0.7281],
        [0.8183, 0.7281],
        [0.8183, 0.7281]])
cand -(colsum_logp)/sum(mask) * m
 tensor([[0.8183, 0.7281],
        [0.8183, 0.0000],
        [0.0000, 0.7281]])
cand with per-col denom colmask
tensor([[1.6367, 1.4563],
        [1.6367, 0.0000],
        [0.0000, 1.4563]])

[stdout]
W (masked)
 tensor([[0.8183, 0.7281],
        [0.8183, 0.7281],
        [0.8183, 0.7281]])
cand -(colsum_logp)/sum(mask) * m
 tensor([[0.8183, 0.7281],
        [0.8183, 0.0000],
        [0.0000, 0.7281]])
cand with per-col denom colmask
tensor([[1.6367, 1.4563],
        [1.6367, 0.0000],
        [0.0000, 1.4563]])

Test full reverse_kl_objective candidate

shell
cat > /app/_p14.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
def cand(logits,ref,labels,mask,adv,beta):
    logp=tlp(logits,labels); reflp=tlp(ref,labels); d=logp-reflp
    kl=torch.exp(-d)+d-1
    M=mask.sum()
    term_adv=-(( (logp*mask).sum(0) )*adv.sum(0)).sum()/M
    # KL: sum over batch of masked_mean over tokens
    seqden=mask.sum(1).clamp(min=1e-9)
    term_kl=beta*((kl*mask).sum(1)/seqden).sum()
    return term_adv+term_kl
for seed in [1,2,3,4]:
    g=torch.Generator().manual_seed(seed)
    B,T,V=torch.randint(2,5,(1,),generator=g).item(),torch.randint(2,5,(1,),generator=g).item(),5
    logits=torch.randn(B,T,V,generator=g);ref=torch.randn(B,T,V,generator=g)
    labels=torch.randint(0,V,(B,T),generator=g)
    mask=(torch.rand(B,T,generator=g)>0.3).float()
    adv=torch.randn(B,T,generator=g); beta=0.1
    r=R.reverse_kl_objective(logits,ref,labels,mask,adv,beta).item()
    c=cand(logits,ref,labels,mask,adv,beta).item()
    print(seed,"ref",round(r,5),"cand",round(c,5),"match",abs(r-c)<1e-4)
EOF
python3 /app/_p14.py
output
1 ref 1.41706 cand 1.45427 match False
2 ref 0.86723 cand 0.80732 match False
3 ref -0.43945 cand -0.41503 match False
4 ref -0.5213 cand -0.55634 match False

[stdout]
1 ref 1.41706 cand 1.45427 match False
2 ref 0.86723 cand 0.80732 match False
3 ref -0.43945 cand -0.41503 match False
4 ref -0.5213 cand -0.55634 match False

Separate adv and kl term mask handling

shell
cat > /app/_p15.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
for seed in [1,2,3]:
    g=torch.Generator().manual_seed(seed)
    B,T,V=3,3,5
    logits=torch.randn(B,T,V,generator=g);ref=torch.randn(B,T,V,generator=g)
    labels=torch.randint(0,V,(B,T),generator=g)
    mask=(torch.rand(B,T,generator=g)>0.3).float()
    adv=torch.randn(B,T,generator=g)
    logp=tlp(logits,labels); reflp=tlp(ref,labels); d=logp-reflp
    kl=torch.exp(-d)+d-1
    M=mask.sum()
    # adv term (beta 0)
    at=R.reverse_kl_objective(logits,ref,labels,mask,adv,0.0).item()
    c1=-(((logp*mask).sum(0))*adv.sum(0)).sum()/M
    c2=-((logp*mask)*adv).sum()/M          # standard masked mean adv*logp
    c3=-((logp*mask)*adv*mask).sum()/M
    print(seed,"ADV ref",round(at,5),"colsum",round(c1.item(),5),"std_mm",round(c2.item(),5))
    # kl term (adv 0)
    kt=R.reverse_kl_objective(logits,ref,labels,mask,torch.zeros_like(adv),0.1).item()
    seqden=mask.sum(1).clamp(min=1e-9)
    k1=0.1*((kl*mask).sum(1)/seqden).sum()
    k2=0.1*((kl*mask).sum()/M)             # global masked mean
    k3=0.1*(kl*mask).sum()                 # masked sum
    print("   KL ref",round(kt,5),"seqmean_sum",round(k1.item(),5),"glob_mm",round(k2.item(),5),"msum",round(k3.item(),5))
EOF
python3 /app/_p15.py
output
1 ADV ref 0.44982 colsum 0.44982 std_mm 0.37927
   KL ref 0.20008 seqmean_sum 0.19831 glob_mm 0.06669 msum 0.46684
2 ADV ref 2.46023 colsum 2.46023 std_mm 0.66411
   KL ref 0.29037 seqmean_sum 0.54263 glob_mm 0.09679 msum 0.58073
3 ADV ref -3.13367 colsum -3.13367 std_mm -0.49742
   KL ref 0.40452 seqmean_sum 0.4237 glob_mm 0.13484 msum 0.94389

[stdout]
1 ADV ref 0.44982 colsum 0.44982 std_mm 0.37927
   KL ref 0.20008 seqmean_sum 0.19831 glob_mm 0.06669 msum 0.46684
2 ADV ref 2.46023 colsum 2.46023 std_mm 0.66411
   KL ref 0.29037 seqmean_sum 0.54263 glob_mm 0.09679 msum 0.58073
3 ADV ref -3.13367 colsum -3.13367 std_mm -0.49742
   KL ref 0.40452 seqmean_sum 0.4237 glob_mm 0.13484 msum 0.94389

Recover KL reduction coefficients

shell
cat > /app/_p16.py << 'EOF'
import torch, reforacle as R
B,T,V=3,3,5
labels=torch.zeros(B,T,dtype=torch.long)
mask=torch.tensor([[1.,1.,0.],[1.,0.,1.],[1.,1.,1.]])
ref=torch.zeros(B,T,V)  # uniform, reflp=-log5
refl=-torch.log(torch.tensor(5.0))
C=torch.zeros(B,T)
for i in range(B):
  for t in range(T):
    logits=torch.zeros(B,T,V)
    logits[i,t,0]=2.0  # change logp at this token
    lp=torch.log_softmax(logits[i,t],-1)[0]
    d=lp-refl
    kl=torch.exp(-d)+d-1
    val=R.reverse_kl_objective(logits,ref,labels,mask,torch.zeros(B,T),1.0).item()
    C[i,t]=val/kl.item() if kl.item()!=0 else 0
print("C (kl coeffs), mask=\n",mask)
print(C)
seqden=mask.sum(1,keepdim=True)
print("mask/seqden\n", mask/seqden)
print("mask/summask\n", mask/mask.sum())
EOF
python3 /app/_p16.py
output
C (kl coeffs), mask=
 tensor([[1., 1., 0.],
        [1., 0., 1.],
        [1., 1., 1.]])
tensor([[0.4286, 0.4286, 0.0000],
        [0.4286, 0.0000, 0.4286],
        [0.4286, 0.4286, 0.4286]])
mask/seqden
 tensor([[0.5000, 0.5000, 0.0000],
        [0.5000, 0.0000, 0.5000],
        [0.3333, 0.3333, 0.3333]])
mask/summask
 tensor([[0.1429, 0.1429, 0.0000],
        [0.1429, 0.0000, 0.1429],
        [0.1429, 0.1429, 0.1429]])

[stdout]
C (kl coeffs), mask=
 tensor([[1., 1., 0.],
        [1., 0., 1.],
        [1., 1., 1.]])
tensor([[0.4286, 0.4286, 0.0000],
        [0.4286, 0.0000, 0.4286],
        [0.4286, 0.4286, 0.4286]])
mask/seqden
 tensor([[0.5000, 0.5000, 0.0000],
        [0.5000, 0.0000, 0.5000],
        [0.3333, 0.3333, 0.3333]])
mask/summask
 tensor([[0.1429, 0.1429, 0.0000],
        [0.1429, 0.0000, 0.1429],
        [0.1429, 0.1429, 0.1429]])

Careful KL term isolation

shell
cat > /app/_p17.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
g=torch.Generator().manual_seed(1)
B,T,V=3,3,5
logits=torch.randn(B,T,V,generator=g);ref=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g)
mask=(torch.rand(B,T,generator=g)>0.3).float()
logp=tlp(logits,labels); reflp=tlp(ref,labels); d=logp-reflp
kl=torch.exp(-d)+d-1
kt=R.reverse_kl_objective(logits,ref,labels,mask,torch.zeros(B,T),1.0).item()
print("mask\n",mask)
print("kl\n",kl)
print("ref kt(beta1)",kt)
print("masked_sum kl", (kl*mask).sum().item())
print("masked_mean kl", ((kl*mask).sum()/mask.sum()).item())
print("sum over batch of masked_mean_t:", ((kl*mask).sum(1)/mask.sum(1).clamp(min=1)).sum().item())
# maybe mean over batch of masked_sum_t
print("colsum coupling like adv? -(masksum_kl col)*ones", None)
EOF
python3 /app/_p17.py
output
mask
 tensor([[1., 1., 1.],
        [1., 1., 1.],
        [1., 0., 0.]])
kl
 tensor([[1.0176, 0.2651, 0.2840],
        [0.3212, 1.3013, 0.8388],
        [0.6404, 0.1944, 1.6581]])
ref kt(beta1) 2.000753164291382
masked_sum kl 4.668424129486084
masked_mean kl 0.6669177412986755
sum over batch of masked_mean_t: 1.9830830097198486
colsum coupling like adv? -(masksum_kl col)*ones None

[stdout]
mask
 tensor([[1., 1., 1.],
        [1., 1., 1.],
        [1., 0., 0.]])
kl
 tensor([[1.0176, 0.2651, 0.2840],
        [0.3212, 1.3013, 0.8388],
        [0.6404, 0.1944, 1.6581]])
ref kt(beta1) 2.000753164291382
masked_sum kl 4.668424129486084
masked_mean kl 0.6669177412986755
sum over batch of masked_mean_t: 1.9830830097198486
colsum coupling like adv? -(masksum_kl col)*ones None

Test reverse_kl_objective full hypothesis

shell
cat > /app/_p18.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
def cand(logits,ref,labels,mask,adv,beta):
    logp=tlp(logits,labels); reflp=tlp(ref,labels); d=logp-reflp
    kl=torch.exp(-d)+d-1
    B=logits.shape[0]; M=mask.sum()
    term_adv=-(((logp*mask).sum(0))*adv.sum(0)).sum()/M
    term_kl=beta*B*((kl*mask).sum()/M)
    return term_adv+term_kl
for seed in range(1,7):
    g=torch.Generator().manual_seed(seed)
    B=torch.randint(2,5,(1,),generator=g).item();T=torch.randint(2,5,(1,),generator=g).item();V=6
    logits=torch.randn(B,T,V,generator=g);ref=torch.randn(B,T,V,generator=g)
    labels=torch.randint(0,V,(B,T),generator=g)
    mask=(torch.rand(B,T,generator=g)>0.3).float()
    adv=torch.randn(B,T,generator=g); beta=torch.rand(1,generator=g).item()
    r=R.reverse_kl_objective(logits,ref,labels,mask,adv,beta).item()
    c=cand(logits,ref,labels,mask,adv,beta).item()
    print(seed,round(r,5),round(c,5),abs(r-c)<1e-4)
EOF
python3 /app/_p18.py
output
1 -5.35465 -5.35465 True
2 0.06041 0.06041 True
3 2.90422 2.90422 True
4 -1.26715 -1.26715 True
5 -2.28249 -2.28249 True
6 7.86388 7.86388 True

[stdout]
1 -5.35465 -5.35465 True
2 0.06041 0.06041 True
3 2.90422 2.90422 True
4 -1.26715 -1.26715 True
5 -2.28249 -2.28249 True
6 7.86388 7.86388 True

Probe reductions, shapes, edge cases

shell
cat > /app/_p19.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
g=torch.Generator().manual_seed(3)
# entropy reduction B=2
logits=torch.randn(2,3,4,generator=g); mask=torch.ones(2,3)
p=logits.softmax(-1); ent=-(p*logits.log_softmax(-1)).sum(-1)
print("entropy ref", R.entropy(logits,mask).item(), "plain mean", ent.mean().item(), "B*mean", 2*ent.mean().item())
# grpo eps with zero-variance group
r=torch.tensor([5.,5.,5.,1.,2.,3.])
print("grpo constgroup scaleT", R.grpo_advantages(r,3,True))
# gae 2d?
try:
    rw=torch.randn(2,3,generator=g); vv=torch.randn(2,3,generator=g); nv=torch.randn(2,generator=g)
    print("gae2d", R.gae(rw,vv,nv,0.9,0.95).shape)
except Exception as e: print("gae2d ERR",e)
# discounted_returns 2d?
try:
    print("discret2d", R.discounted_returns(torch.randn(2,3,generator=g),0.9).shape)
except Exception as e: print("dr2d ERR",e)
# cross_entropy all-ignored
print("ce all ignore", R.cross_entropy(torch.randn(2,4,generator=g),torch.tensor([-100,-100]),-100))
# top_k ties
print("top_k tie", R.top_k_mask(torch.tensor([[1.,1.,1.,0.]]),2))
EOF
python3 /app/_p19.py
output
entropy ref 1.22975754737854 plain mean 1.22975754737854 B*mean 2.45951509475708
grpo constgroup scaleT tensor([ 0.0000, 0.0000, 0.0000, -1.2247, 0.0000, 1.2247])
gae2d ERR The size of tensor a (3) must match the size of tensor b (2) at non-singleton dimension 0
discret2d torch.Size([2, 3])
ce all ignore tensor(0.)
top_k tie tensor([[ True, True, True, False]])

[stdout]
entropy ref 1.22975754737854 plain mean 1.22975754737854 B*mean 2.45951509475708
grpo constgroup scaleT tensor([ 0.0000, 0.0000, 0.0000, -1.2247, 0.0000, 1.2247])
gae2d ERR The size of tensor a (3) must match the size of tensor b (2) at non-singleton dimension 0
discret2d torch.Size([2, 3])
ce all ignore tensor(0.)
top_k tie tensor([[ True, True, True, False]])

grpo eps and discounted_returns axis

shell
cat > /app/_p20.py << 'EOF'
import torch, reforacle as R
# grpo eps: small std group
r=torch.tensor([0.,0.0001,0.,0.,0.,0.0])  # group1 tiny std
out=R.grpo_advantages(r,3,True)
print("grpo tiny", out)
g1=r[:3]; c=g1-g1.mean(); sb=g1.std(unbiased=False)
for eps in [0,1e-8,1e-6,1e-5,1e-4]:
    print("eps",eps,(c/(sb+eps)).tolist())
# discounted_returns axis for 2d
x=torch.tensor([[1.,2.,3.],[10.,20.,30.]])
print("dr2d\n", R.discounted_returns(x,0.5))
EOF
python3 /app/_p20.py
output
grpo tiny tensor([-0.6924, 1.3848, -0.6924, 0.0000, 0.0000, 0.0000])
eps 0 [-0.7071068286895752, 1.4142135381698608, -0.7071068286895752]
eps 1e-08 [-0.7069568634033203, 1.413913607597351, -0.7069568634033203]
eps 1e-06 [-0.6924183964729309, 1.3848366737365723, -0.6924183964729309]
eps 1e-05 [-0.5833579301834106, 1.1667157411575317, -0.5833579301834106]
eps 0.0001 [-0.2265409380197525, 0.4530818462371826, -0.2265409380197525]
dr2d
 tensor([[ 6., 12., 18.],
        [10., 20., 30.]])

[stdout]
grpo tiny tensor([-0.6924, 1.3848, -0.6924, 0.0000, 0.0000, 0.0000])
eps 0 [-0.7071068286895752, 1.4142135381698608, -0.7071068286895752]
eps 1e-08 [-0.7069568634033203, 1.413913607597351, -0.7069568634033203]
eps 1e-06 [-0.6924183964729309, 1.3848366737365723, -0.6924183964729309]
eps 1e-05 [-0.5833579301834106, 1.1667157411575317, -0.5833579301834106]
eps 0.0001 [-0.2265409380197525, 0.4530818462371826, -0.2265409380197525]
dr2d
 tensor([[ 6., 12., 18.],
        [10., 20., 30.]])

Probe dpo_sequence_loss

shell
cat > /app/_p21.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
def seqlp(lg,lb,m,norm):
    t=(tlp(lg,lb)*m).sum(-1)
    return t/m.sum(-1) if norm else t
g=torch.Generator().manual_seed(11)
B,T,V=2,4,5
pc_l=torch.randn(B,T,V,generator=g);pr_l=torch.randn(B,T,V,generator=g)
rc_l=torch.randn(B,T,V,generator=g);rr_l=torch.randn(B,T,V,generator=g)
cl=torch.randint(0,V,(B,T),generator=g);rl=torch.randint(0,V,(B,T),generator=g)
cm=(torch.rand(B,T,generator=g)>0.3).float();rm=(torch.rand(B,T,generator=g)>0.3).float()
beta=0.1;ls=0.0
ref=R.dpo_sequence_loss(pc_l,pr_l,rc_l,rr_l,cl,rl,cm,rm,beta,ls)
print("ref",ref.item())
for norm in [False,True]:
    pc=seqlp(pc_l,cl,cm,norm);pr=seqlp(pr_l,rl,rm,norm)
    rc=seqlp(rc_l,cl,cm,norm);rr=seqlp(rr_l,rl,rm,norm)
    logits=beta*((pc-rc)-(pr-rr))
    loss=(-(1-ls)*torch.nn.functional.logsigmoid(logits)-ls*torch.nn.functional.logsigmoid(-logits)).mean()
    print("norm",norm,loss.item())
EOF
python3 /app/_p21.py
output
ref 0.7345190048217773
norm False 0.7345190048217773
norm True 0.7023628950119019

[stdout]
ref 0.7345190048217773
norm False 0.7345190048217773
norm True 0.7023628950119019

Probe rloo_objective

shell
cat > /app/_p22.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
g=torch.Generator().manual_seed(13)
B,T,V=6,3,5; gs=3
logits=torch.randn(B,T,V,generator=g);old=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g)
mask=(torch.rand(B,T,generator=g)>0.3).float()
rewards=torch.randn(B,generator=g)
cl,ch=0.2,0.3
ref=R.rloo_objective(logits,old,labels,mask,rewards,gs,cl,ch)
print("ref",ref.item())
# rloo advantages
adv=R.rloo_advantages(rewards,gs)  # [B]
A=adv.unsqueeze(1).expand(B,T)
logp=tlp(logits,labels);oldlp=tlp(old,labels)
ratio=torch.exp(logp-oldlp)
surr1=ratio*A; surr2=torch.clamp(ratio,1-cl,1+ch)*A
pg=-torch.min(surr1,surr2)
print("plain mm", ((pg*mask).sum()/mask.sum()).item())
print("B*mm", B*((pg*mask).sum()/mask.sum()).item())
print("sum-batch of seqmean", ((pg*mask).sum(1)/mask.sum(1).clamp(min=1)).sum().item())
EOF
python3 /app/_p22.py
output
ref 0.29768627882003784
plain mm 0.29768627882003784
B*mm 1.786117672920227
sum-batch of seqmean 1.788561224937439

[stdout]
ref 0.29768627882003784
plain mm 0.29768627882003784
B*mm 1.786117672920227
sum-batch of seqmean 1.788561224937439

Probe grpo_objective

shell
cat > /app/_p23.py << 'EOF'
import torch, reforacle as R
def tlp(lg,lb): return torch.log_softmax(lg,-1).gather(-1,lb.unsqueeze(-1)).squeeze(-1)
g=torch.Generator().manual_seed(17)
B,T,V=6,3,5; gs=3
logits=torch.randn(B,T,V,generator=g);old=torch.randn(B,T,V,generator=g);refl=torch.randn(B,T,V,generator=g)
labels=torch.randint(0,V,(B,T),generator=g)
mask=(torch.rand(B,T,generator=g)>0.3).float()
rewards=torch.randn(B,generator=g)
beta,cl,ch=0.05,0.2,0.2
def klpen(d,est):
    if est=="k1": return d
    if est=="k2": return 0.5*d*d
    if est=="k3": return torch.exp(-d)+d-1
for est in ["k1","k2","k3"]:
  ref=R.grpo_objective(logits,old,refl,labels,mask,rewards,gs,beta,cl,ch,True,est).item()
  adv=R.grpo_advantages(rewards,gs,True).unsqueeze(1).expand(B,T)
  logp=tlp(logits,labels);oldlp=tlp(old,labels);reflp=tlp(refl,labels)
  ratio=torch.exp(logp-oldlp)
  pg=-torch.min(ratio*adv,torch.clamp(ratio,1-cl,1+ch)*adv)
  kl=klpen(logp-reflp,est)
  per=pg+beta*kl
  mm=(per*mask).sum()/mask.sum()
  print(est,"ref",round(ref,5),"pg+bkl mm",round(mm.item(),5),"match",abs(ref-mm.item())<1e-5)
EOF
python3 /app/_p23.py
output
k1 ref 0.41223 pg+bkl mm 0.41223 match True
k2 ref 0.46354 pg+bkl mm 0.46354 match True
k3 ref 0.72606 pg+bkl mm 0.72606 match True

[stdout]
k1 ref 0.41223 pg+bkl mm 0.41223 match True
k2 ref 0.46354 pg+bkl mm 0.46354 match True
k3 ref 0.72606 pg+bkl mm 0.72606 match True

Probe ppo_objective

shell
cat > /app/_p24.py << 'EOF'
import torch, reforacle as R
g=torch.Generator().manual_seed(19)
T=5
rewards=torch.randn(T,generator=g);values=torch.randn(T,generator=g);old_values=torch.randn(T,generator=g)
logp=torch.randn(T,generator=g)*0.1;old_logp=torch.randn(T,generator=g)*0.1
nv=torch.randn(1,generator=g)[0]
gamma,lam,cl,ch,vfc,vfcoef=0.99,0.95,0.2,0.2,0.2,0.5
ref=R.ppo_objective(rewards,values,old_values,logp,old_logp,nv,gamma,lam,cl,ch,vfc,vfcoef).item()
print("ref",ref)
gae=R.gae(rewards,values,nv,gamma,lam)
returns=gae+values
ratio=torch.exp(logp-old_logp)
def pgloss(adv):
    return torch.max(-adv*ratio,-adv*torch.clamp(ratio,1-cl,1+ch)).mean()
vl=R.value_loss(values,old_values,returns,vfc)
for label,adv in [("raw",gae),("whiten",(gae-gae.mean())/(gae.std(unbiased=False)+1e-8)),
                  ("whiten_ub",(gae-gae.mean())/(gae.std(unbiased=True)+1e-8))]:
    tot=pgloss(adv)+vfcoef*vl
    print(label,"pg",round(pgloss(adv).item(),5),"tot",round(tot.item(),5),"match",abs(tot.item()-ref)<1e-5)
print("vl",vl.item())
EOF
python3 /app/_p24.py
output
ref 1.6760817766189575
raw pg 0.55052 tot 2.191 match False
whiten pg 0.0356 tot 1.67608 match True
whiten_ub pg 0.03184 tot 1.67232 match False
vl 3.2809576988220215

[stdout]
ref 1.6760817766189575
raw pg 0.55052 tot 2.191 match False
whiten pg 0.0356 tot 1.67608 match True
whiten_ub pg 0.03184 tot 1.67232 match False
vl 3.2809576988220215

Final edge conventions

shell
cat > /app/_p25.py << 'EOF'
import torch, reforacle as R
# normalize denom: std+eps vs sqrt(var+eps), big eps
x=torch.tensor([1.,2.,3.,4.])
print("norm eps1", R.normalize(x,1.0).tolist())
print(" (x-m)/(std_ub+1)", ((x-x.mean())/(x.std(unbiased=True)+1)).tolist())
print(" (x-m)/sqrt(var_ub+1)", ((x-x.mean())/torch.sqrt(x.var(unbiased=True)+1)).tolist())
# masked_max 2d
v=torch.tensor([[1.,5.,2.],[8.,3.,9.]]);m=torch.tensor([[1.,0.,1.],[0.,1.,1.]])
print("masked_max2d", R.masked_max(v,m))
# mode
print("mode 1122", R.mode_label(torch.tensor([1,1,2,2])), "torch.mode", torch.mode(torch.tensor([1,1,2,2])).values)
# smoothed_nll
logits=torch.tensor([[2.,1.,0.,-1.]]);labels=torch.tensor([0])
lp=logits.log_softmax(-1); nll=-lp[0,0]; K=4
s=0.1
print("smoothed ref", R.smoothed_nll(logits,labels,s).item())
print(" cand (1-s)*nll - s*mean(lp)", ((1-s)*nll - s*lp.mean()).item())
# top_p exact boundary
print("top_p 0.9", R.top_p_mask(torch.tensor([[0.5,0.3,0.15,0.05]]),0.9))
EOF
python3 /app/_p25.py
output
norm eps1 [-0.6547375321388245, -0.21824583411216736, 0.21824583411216736, 0.6547375321388245]
 (x-m)/(std_ub+1) [-0.6547375321388245, -0.21824583411216736, 0.21824583411216736, 0.6547375321388245]
 (x-m)/sqrt(var_ub+1) [-0.9185587167739868, -0.3061862289905548, 0.3061862289905548, 0.9185587167739868]
masked_max2d tensor(9.)
mode 1122 tensor(2) torch.mode tensor(1)
smoothed ref 0.5901897549629211
 cand (1-s)*nll - s*mean(lp) 0.5901897549629211
top_p 0.9 tensor([[0.5263, 0.3158, 0.1579, 0.0000]])

[stdout]
norm eps1 [-0.6547375321388245, -0.21824583411216736, 0.21824583411216736, 0.6547375321388245]
 (x-m)/(std_ub+1) [-0.6547375321388245, -0.21824583411216736, 0.21824583411216736, 0.6547375321388245]
 (x-m)/sqrt(var_ub+1) [-0.9185587167739868, -0.3061862289905548, 0.3061862289905548, 0.9185587167739868]
masked_max2d tensor(9.)
mode 1122 tensor(2) torch.mode tensor(1)
smoothed ref 0.5901897549629211
 cand (1-s)*nll - s*mean(lp) 0.5901897549629211
top_p 0.9 tensor([[0.5263, 0.3158, 0.1579, 0.0000]])
write · /app/objectives.py
from __future__ import annotations

import torch


# ---------------------------------------------------------------------------
# helpers
# ---------------------------------------------------------------------------
def _token_logprobs(logits, labels):
    lp = torch.log_softmax(logits, dim=-1)
    return lp.gather(-1, labels.unsqueeze(-1)).squeeze(-1)


def _masked_mean(values, mask, axis=None):
    if axis is None:
        num = (values * mask).sum()
        den = mask.sum()
    else:
        num = (values * mask).sum(axis)
        den = mask.sum(axis)
    return num / den.clamp(min=1e-12)


# ---------------------------------------------------------------------------
# primitives
# ---------------------------------------------------------------------------
def masked_mean(values, mask, axis=None):
    return _masked_mean(values, mask, axis)


def masked_sum(values, mask, axis=None):
    if axis is None:
        return (values * mask).sum()
    return (values * mask).sum(axis)


def logsumexp(x, axis):
    return torch.logsumexp(x, dim=axis)


def log_softmax(x, axis):
    return torch.log_softmax(x, dim=axis)


def selective_logprobs(logits, labels, mask):
    return (_token_logprobs(logits, labels) * mask).sum(-1)


def token_logprobs(logits, labels):
    return _token_logprobs(logits, labels)


def sequence_logprob(logits, labels, mask, length_normalize):
    lp = _token_logprobs(logits, labels)
    s = (lp * mask).sum(-1)
    if length_normalize:
        return s / mask.sum(-1).clamp(min=1e-12)
    return s


def entropy(logits, mask):
    lp = torch.log_softmax(logits, dim=-1)
    p = lp.exp()
    ent = -(p * lp).sum(-1)
    return _masked_mean(ent, mask)


def dpo_loss(pc, pr, rc, rr, beta, label_smoothing):
    logits = beta * ((pc - rc) - (pr - rr))
    loss = (-(1.0 - label_smoothing) * torch.nn.functional.logsigmoid(logits)
            - label_smoothing * torch.nn.functional.logsigmoid(-logits))
    return loss.mean()


def ipo_loss(pc, pr, rc, rr, beta):
    h = (pc - rc) - (pr - rr)
    loss = (h - 1.0 / (2.0 * beta)) ** 2
    return loss.mean()


def grpo_advantages(rewards, group_size, scale_by_std):
    g = rewards.view(-1, group_size)
    mean = g.mean(dim=1, keepdim=True)
    centered = g - mean
    if scale_by_std:
        std = g.std(dim=1, unbiased=False, keepdim=True)
        centered = centered / (std + 1e-6)
    return centered.reshape(rewards.shape)


def gae(rewards, values, next_value, gamma, lam):
    T = rewards.shape[0]
    adv = torch.zeros_like(rewards)
    last = torch.zeros_like(next_value)
    for t in range(T - 1, -1, -1):
        v_next = next_value if t == T - 1 else values[t + 1]
        delta = rewards[t] + gamma * v_next - values[t]
        last = delta + gamma * lam * last
        adv[t] = last
    return adv


def kl_penalty(logp, ref_logp, estimator):
    d = logp - ref_logp
    if estimator == "k1":
        return d
    if estimator == "k2":
        return 0.5 * d * d
    # k3
    return torch.exp(-d) + d - 1.0


def clipped_pg_loss(logp, old_logp, advantages, mask, clip_low, clip_high):
    ratio = torch.exp(logp - old_logp)
    unclipped = -advantages * ratio
    clipped = -advantages * torch.clamp(ratio, 1.0 - clip_low, 1.0 + clip_high)
    loss = torch.maximum(unclipped, clipped)
    return _masked_mean(loss, mask)


def value_loss(values, old_values, returns, clip):
    v_clipped = old_values + torch.clamp(values - old_values, -clip, clip)
    l1 = (values - returns) ** 2
    l2 = (v_clipped - returns) ** 2
    return 0.5 * torch.maximum(l1, l2).mean()


def whiten(values, mask, shift_mean):
    mean = _masked_mean(values, mask)
    var = _masked_mean((values - mean) ** 2, mask)
    whitened = (values - mean) / (torch.sqrt(var) + 1e-8)
    if not shift_mean:
        whitened = whitened + mean
    return whitened


def dpo_sequence_loss(pc_logits, pr_logits, rc_logits, rr_logits,
                      chosen_labels, rejected_labels, chosen_mask, rejected_mask,
                      beta, label_smoothing):
    pc = (_token_logprobs(pc_logits, chosen_labels) * chosen_mask).sum(-1)
    pr = (_token_logprobs(pr_logits, rejected_labels) * rejected_mask).sum(-1)
    rc = (_token_logprobs(rc_logits, chosen_labels) * chosen_mask).sum(-1)
    rr = (_token_logprobs(rr_logits, rejected_labels) * rejected_mask).sum(-1)
    return dpo_loss(pc, pr, rc, rr, beta, label_smoothing)


def grpo_objective(logits, old_logits, ref_logits, labels, completion_mask,
                   rewards, group_size, beta, clip_low, clip_high, scale_by_std,
                   kl_estimator):
    adv = grpo_advantages(rewards, group_size, scale_by_std)
    adv = adv.unsqueeze(-1).expand_as(labels.float() if labels.dtype.is_floating_point
                                      else labels)
    adv = adv.to(logits.dtype)
    logp = _token_logprobs(logits, labels)
    old_logp = _token_logprobs(old_logits, labels)
    ref_logp = _token_logprobs(ref_logits, labels)
    ratio = torch.exp(logp - old_logp)
    unclipped = -adv * ratio
    clipped = -adv * torch.clamp(ratio, 1.0 - clip_low, 1.0 + clip_high)
    pg = torch.maximum(unclipped, clipped)
    kl = kl_penalty(logp, ref_logp, kl_estimator)
    return _masked_mean(pg + beta * kl, completion_mask)


def ppo_objective(rewards, values, old_values, logp, old_logp, next_value,
                  gamma, lam, clip_low, clip_high, vf_clip, vf_coef):
    adv = gae(rewards, values, next_value, gamma, lam)
    returns = adv + values
    mean = adv.mean()
    var = ((adv - mean) ** 2).mean()
    adv_w = (adv - mean) / (torch.sqrt(var) + 1e-8)
    ratio = torch.exp(logp - old_logp)
    unclipped = -adv_w * ratio
    clipped = -adv_w * torch.clamp(ratio, 1.0 - clip_low, 1.0 + clip_high)
    pg = torch.maximum(unclipped, clipped).mean()
    vl = value_loss(values, old_values, returns, vf_clip)
    return pg + vf_coef * vl


def rloo_advantages(rewards, group_size):
    g = rewards.view(-1, group_size)
    total = g.sum(dim=1, keepdim=True)
    baseline = (total - g) / (group_size - 1)
    return (g - baseline).reshape(rewards.shape)


def reverse_kl(logp, ref_logp):
    d = logp - ref_logp
    return torch.exp(-d) + d - 1.0


def importance_ratio(logp, old_logp, clip):
    ratio = torch.exp(logp - old_logp)
    if clip is not None:
        ratio = torch.clamp(ratio, 1.0 - clip, 1.0 + clip)
    return ratio


def discounted_returns(rewards, gamma):
    T = rewards.shape[0]
    out = torch.zeros_like(rewards)
    last = torch.zeros_like(rewards[0])
    for t in range(T - 1, -1, -1):
        last = rewards[t] + gamma * last
        out[t] = last
    return out


def normalize(x, eps):
    mean = x.mean()
    std = x.std(unbiased=True)
    return (x - mean) / (std + eps)


def top_p_mask(probs, p):
    sorted_probs, sorted_idx = torch.sort(probs, dim=-1, descending=True)
    cumsum = sorted_probs.cumsum(dim=-1)
    remove = cumsum >= p
    # shift right so the token crossing the threshold is kept
    remove[..., 1:] = remove[..., :-1].clone()
    remove[..., 0] = False
    sorted_probs = sorted_probs.masked_fill(remove, 0.0)
    kept = torch.zeros_like(probs)
    kept.scatter_(-1, sorted_idx, sorted_probs)
    kept = kept / kept.sum(dim=-1, keepdim=True).clamp(min=1e-12)
    return kept


def smoothed_nll(logits, labels, smoothing):
    lp = torch.log_softmax(logits, dim=-1)
    nll = -lp.gather(-1, labels.unsqueeze(-1)).squeeze(-1)
    smooth = -lp.mean(dim=-1)
    loss = (1.0 - smoothing) * nll + smoothing * smooth
    return loss.mean()


def bradley_terry_logit(chosen_reward, rejected_reward, beta):
    return beta * (chosen_reward - rejected_reward)


def rloo_objective(logits, old_logits, labels, mask, rewards, group_size, clip_low, clip_high):
    adv = rloo_advantages(rewards, group_size)
    adv = adv.unsqueeze(-1).to(logits.dtype)
    logp = _token_logprobs(logits, labels)
    old_logp = _token_logprobs(old_logits, labels)
    ratio = torch.exp(logp - old_logp)
    unclipped = -adv * ratio
    clipped = -adv * torch.clamp(ratio, 1.0 - clip_low, 1.0 + clip_high)
    pg = torch.maximum(unclipped, clipped)
    return _masked_mean(pg, mask)


def reverse_kl_objective(logits, ref_logits, labels, mask, advantages, beta):
    B = logits.shape[0]
    logp = _token_logprobs(logits, labels)
    ref_logp = _token_logprobs(ref_logits, labels)
    d = logp - ref_logp
    kl = torch.exp(-d) + d - 1.0
    M = mask.sum().clamp(min=1e-12)
    term_adv = -(((logp * mask).sum(0)) * advantages.sum(0)).sum() / M
    term_kl = beta * B * ((kl * mask).sum() / M)
    return term_adv + term_kl


def cross_entropy(logits, labels, ignore_index):
    lp = torch.log_softmax(logits, dim=-1)
    valid = labels != ignore_index
    safe = labels.clamp(min=0)
    nll = -lp.gather(-1, safe.unsqueeze(-1)).squeeze(-1)
    nll = nll * valid
    return nll.sum() / valid.sum().clamp(min=1e-12)


def top_k_mask(logits, k):
    kth = torch.topk(logits, k, dim=-1).values[..., -1:]
    return logits >= kth


def group_mean_baseline(rewards, group_size):
    g = rewards.view(-1, group_size)
    mean = g.mean(dim=1, keepdim=True)
    return (g - mean).reshape(rewards.shape)


def lambda_returns(rewards, values, next_value, gamma, lam):
    return gae(rewards, values, next_value, gamma, lam) + values


def symmetric_kl(logp, ref_logp):
    d = logp - ref_logp
    rev = torch.exp(-d) + d - 1.0
    fwd = torch.exp(d) - d - 1.0
    return 0.5 * (rev + fwd)


def huber_value_loss(values, returns, delta):
    err = values - returns
    a = err.abs()
    loss = torch.where(a <= delta, 0.5 * err * err, delta * (a - 0.5 * delta))
    return loss.mean()


def normalized_entropy(logits, mask):
    V = logits.shape[-1]
    return entropy(logits, mask) / torch.log(torch.tensor(float(V)))


def clip_fraction(logp, old_logp, clip):
    ratio = torch.exp(logp - old_logp)
    clipped = (ratio < 1.0 - clip) | (ratio > 1.0 + clip)
    return clipped.float().mean()


def masked_whiten(values, mask, shift_mean):
    mean = _masked_mean(values, mask)
    var = _masked_mean((values - mean) ** 2, mask)
    whitened = (values - mean) / (torch.sqrt(var) + 1e-8)
    whitened = whitened * mask
    if not shift_mean:
        whitened = whitened + mean
    return whitened


def logprob_at_temperature(logits, labels, temperature):
    scaled = logits / temperature
    return _token_logprobs(scaled, labels)


def advantage_mean_std(advantages, mask):
    mean = _masked_mean(advantages, mask)
    var = _masked_mean((advantages - mean) ** 2, mask)
    std = torch.sqrt(var)
    return torch.stack([mean, std])


def argmax_tokens(logits):
    return logits.argmax(dim=-1)


def mode_label(labels):
    vals, counts = torch.unique(labels, return_counts=True)
    maxc = counts.max()
    return vals[counts == maxc].max()


def median_reward(rewards):
    return torch.quantile(rewards.float(), 0.5)


def masked_max(values, mask):
    filled = values.masked_fill(mask == 0, float("-inf"))
    return filled.max()


def quantile_lower(x, q):
    return torch.quantile(x.float(), q, interpolation="lower")


def pad_mask_from_lengths(lengths, max_len):
    ar = torch.arange(max_len, device=lengths.device)
    return (ar.unsqueeze(0) < lengths.unsqueeze(1)).float()


def first_nonzero_index(mask):
    nz = mask != 0
    has = nz.any(dim=-1)
    first = nz.float().argmax(dim=-1)
    return torch.where(has, first, torch.full_like(first, -1))


def cumulative_max(x):
    inclusive = torch.cummax(x, dim=0).values
    return torch.cat([x[:1], inclusive[:-1]], dim=0)


def bucketize_reward(rewards, boundaries):
    return torch.bucketize(rewards, boundaries, right=True)

Run full comparison harness

shell
cat > /app/_harness.py << 'EOF'
import torch, reforacle as R
import importlib, objectives
importlib.reload(objectives)
O=objectives
g=torch.Generator().manual_seed(1234)
def rn(*s): return torch.randn(*s,generator=g)
def rint(hi,*s): return torch.randint(0,hi,s,generator=g)
def rmask(*s): return (torch.rand(*s,generator=g)>0.35).float()

results={}
def chk(name,args_fn,n=6,composite=False):
    atol=1e-5 if composite else 1e-6; rtol=1e-4 if composite else 1e-5
    ok=True; err=None
    for _ in range(n):
        args=args_fn()
        try:
            r=getattr(R,name)(*args); m=getattr(O,name)(*args)
        except Exception as e:
            ok=False; err=f"EXC {e}"; break
        r=r if torch.is_tensor(r) else torch.tensor(r)
        m=m if torch.is_tensor(m) else torch.tensor(m)
        if r.shape!=m.shape: ok=False; err=f"shape {r.shape} vs {m.shape}"; break
        if r.dtype!=m.dtype and not (r.is_floating_point() and m.is_floating_point()):
            ok=False; err=f"dtype {r.dtype} vs {m.dtype}"; break
        if not torch.allclose(r.float(),m.float(),atol=atol,rtol=rtol,equal_nan=True):
            ok=False; err=f"val diff max {(r.float()-m.float()).abs().max().item():.2e}"; break
    results[name]=(ok,err)
    print(("PASS" if ok else "FAIL"),name, "" if ok else err)

BT=lambda:(torch.randint(2,5,(1,),generator=g).item(),torch.randint(2,5,(1,),generator=g).item())
def logits_labels_mask():
    B,T=BT(); V=torch.randint(3,7,(1,),generator=g).item()
    return rn(B,T,V),rint(V,B,T),rmask(B,T)

chk("masked_mean",lambda:(rn(4,5),rmask(4,5),None))
chk("masked_mean",lambda:(rn(4,5),rmask(4,5),1))
chk("masked_sum",lambda:(rn(4,5),rmask(4,5),0))
chk("logsumexp",lambda:(rn(3,5),1))
chk("log_softmax",lambda:(rn(3,5),-1))
chk("selective_logprobs",logits_labels_mask)
chk("token_logprobs",lambda:(lambda x:(x[0],x[1]))(logits_labels_mask()))
chk("sequence_logprob",lambda:(*logits_labels_mask(),bool(torch.randint(0,2,(1,),generator=g).item())))
chk("entropy",logits_labels_mask)
def dpo_args():
    B=3; return rn(B),rn(B),rn(B),rn(B),torch.rand(1,generator=g).item()+0.05,torch.rand(1,generator=g).item()*0.2
chk("dpo_loss",dpo_args,composite=True)
chk("ipo_loss",lambda:(rn(3),rn(3),rn(3),rn(3),torch.rand(1,generator=g).item()+0.05),composite=True)
chk("grpo_advantages",lambda:(rn(12),3,bool(torch.randint(0,2,(1,),generator=g).item())))
chk("gae",lambda:(rn(6),rn(6),rn(1)[0],0.99,0.95))
chk("kl_penalty",lambda:(rn(8),rn(8),["k1","k2","k3"][torch.randint(0,3,(1,),generator=g).item()]))
chk("clipped_pg_loss",lambda:(rn(4,5),rn(4,5),rn(4,5),rmask(4,5),0.2,0.3),composite=True)
chk("value_loss",lambda:(rn(6),rn(6),rn(6),0.2),composite=True)
chk("whiten",lambda:(rn(7),rmask(7),bool(torch.randint(0,2,(1,),generator=g).item())))
def dpo_seq_args():
    B,T=3,4; V=5
    return (rn(B,T,V),rn(B,T,V),rn(B,T,V),rn(B,T,V),rint(V,B,T),rint(V,B,T),rmask(B,T),rmask(B,T),0.1,0.05)
chk("dpo_sequence_loss",dpo_seq_args,composite=True)
def grpo_obj_args():
    B,T,V=6,3,5
    return (rn(B,T,V),rn(B,T,V),rn(B,T,V),rint(V,B,T),rmask(B,T),rn(B),3,0.05,0.2,0.2,
            bool(torch.randint(0,2,(1,),generator=g).item()),["k1","k2","k3"][torch.randint(0,3,(1,),generator=g).item()])
chk("grpo_objective",grpo_obj_args,composite=True)
chk("ppo_objective",lambda:(rn(6),rn(6),rn(6),rn(6)*0.1,rn(6)*0.1,rn(1)[0],0.99,0.95,0.2,0.2,0.2,0.5),composite=True)
chk("rloo_advantages",lambda:(rn(12),4))
chk("reverse_kl",lambda:(rn(8),rn(8)))
chk("importance_ratio",lambda:(rn(6)*0.2,rn(6)*0.2,[None,0.2][torch.randint(0,2,(1,),generator=g).item()]))
chk("discounted_returns",lambda:(rn(7),0.9))
chk("normalize",lambda:(rn(8),1e-8))
chk("top_p_mask",lambda:(torch.softmax(rn(4,6),-1),0.7))
chk("smoothed_nll",lambda:(rn(5,7),rint(7,5),0.1),composite=True)
chk("bradley_terry_logit",lambda:(rn(4),rn(4),0.5))
def rloo_obj_args():
    B,T,V=6,3,5
    return (rn(B,T,V),rn(B,T,V),rint(V,B,T),rmask(B,T),rn(B),3,0.2,0.3)
chk("rloo_objective",rloo_obj_args,composite=True)
def rkl_obj_args():
    B,T,V=4,3,5
    return (rn(B,T,V),rn(B,T,V),rint(V,B,T),rmask(B,T),rn(B,T),0.1)
chk("reverse_kl_objective",rkl_obj_args,composite=True)
chk("cross_entropy",lambda:(rn(5,7),rint(7,5),-100))
chk("top_k_mask",lambda:(rn(4,7),2))
chk("group_mean_baseline",lambda:(rn(12),4))
chk("lambda_returns",lambda:(rn(6),rn(6),rn(1)[0],0.99,0.95))
chk("symmetric_kl",lambda:(rn(8),rn(8)))
chk("huber_value_loss",lambda:(rn(6),rn(6),1.0),composite=True)
chk("normalized_entropy",logits_labels_mask)
chk("clip_fraction",lambda:(rn(8)*0.3,rn(8)*0.3,0.2))
chk("masked_whiten",lambda:(rn(7),rmask(7),bool(torch.randint(0,2,(1,),generator=g).item())))
chk("logprob_at_temperature",lambda:(lambda x:(x[0],x[1],1.5))(logits_labels_mask()))
chk("advantage_mean_std",lambda:(rn(4,5),rmask(4,5)))
chk("argmax_tokens",lambda:(rn(3,4,6),))
chk("mode_label",lambda:(rint(4,10),))
chk("median_reward",lambda:(rn(7),))
chk("masked_max",lambda:(rn(4,5),rmask(4,5)))
chk("quantile_lower",lambda:(rn(9),0.3))
chk("pad_mask_from_lengths",lambda:(torch.randint(1,5,(4,),generator=g),5))
chk("first_nonzero_index",lambda:(rmask(4,6),))
chk("cumulative_max",lambda:(rn(6),))
chk("bucketize_reward",lambda:(rn(8),torch.tensor([-1.,0.,1.])))

print("\n== SUMMARY ==")
fails=[k for k,(ok,_) in results.items() if not ok]
print(f"{sum(1 for _,(ok,_) in results.items() if ok)}/{len(results)} passing")
if fails: print("FAILS:",fails)
EOF
python3 /app/_harness.py
output
PASS masked_mean 
PASS masked_mean 
PASS masked_sum 
PASS logsumexp 
PASS log_softmax 
PASS selective_logprobs 
PASS token_logprobs 
PASS sequence_logprob 
FAIL entropy EXC entropy() takes exactly 2 positional arguments (3 given)
PASS dpo_loss 
PASS ipo_loss 
PASS grpo_advantages 
PASS gae 
PASS kl_penalty 
PASS clipped_pg_loss 
PASS value_loss 
FAIL whiten val diff max 2.07e+08
PASS dpo_sequence_loss 
PASS grpo_objective 
PASS ppo_objective 
PASS rloo_advantages 
PASS reverse_kl 
PASS importance_ratio 
PASS discounted_returns 
PASS normalize 
PASS top_p_mask 
PASS smoothed_nll 
PASS bradley_terry_logit 
PASS rloo_objective 
PASS reverse_kl_objective 
PASS cross_entropy 
PASS top_k_mask 
PASS group_mean_baseline 
PASS lambda_returns 
PASS symmetric_kl 
PASS huber_value_loss 
FAIL normalized_entropy EXC normalized_entropy() takes exactly 2 positional arguments (3 given)
PASS clip_fraction 
PASS masked_whiten 
PASS logprob_at_temperature 
PASS advantage_mean_std 
PASS argmax_tokens 
PASS mode_label 
PASS median_reward 
PASS masked_max 
PASS quantile_lower 
PASS pad_mask_from_lengths 
PASS first_nonzero_index 
PASS cumulative_max 
PASS bucketize_reward 

== SUMMARY ==
46/49 passing
FAILS: ['entropy', 'whiten', 'normalized_entropy']

[stdout]
PASS masked_mean 
PASS masked_mean 
PASS masked_sum 
PASS logsumexp 
PASS log_softmax 
PASS selective_logprobs 
PASS token_logprobs 
PASS sequence_logprob 
FAIL entropy EXC entropy() takes exactly 2 positional arguments (3 given)
PASS dpo_loss 
PASS ipo_loss 
PASS grpo_advantages 
PASS gae 
PASS kl_penalty 
PASS clipped_pg_loss 
PASS value_loss 
FAIL whiten val diff max 2.07e+08
PASS dpo_sequence_loss 
PASS grpo_objective 
PASS ppo_objective 
PASS rloo_advantages 
PASS reverse_kl 
PASS importance_ratio 
PASS discounted_returns 
PASS normalize 
PASS top_p_mask 
PASS smoothed_nll 
PASS bradley_terry_logit 
PASS rloo_objective 
PASS reverse_kl_objective 
PASS cross_entropy 
PASS top_k_mask 
PASS group_mean_baseline 
PASS lambda_returns 
PASS symmetric_kl 
PASS huber_value_loss 
FAIL normalized_entropy EXC normalized_entropy() takes exactly 2 positional arguments (3 given)
PASS clip_fraction 
PASS masked_whiten 
PASS logprob_at_temperature 
PASS advantage_mean_std 
PASS argmax_tokens 
PASS mode_label 
PASS median_reward 
PASS masked_max 
PASS quantile_lower 
PASS pad_mask_from_lengths 
PASS first_nonzero_index 
PASS cumulative_max 
PASS bucketize_reward 

== SUMMARY ==
46/49 passing
FAILS: ['entropy', 'whiten', 'normalized_entropy']

Find whiten failing case (empty mask)

shell
cat > /app/_dbg.py << 'EOF'
import torch, reforacle as R, objectives as O
g=torch.Generator().manual_seed(1234)
def rn(*s): return torch.randn(*s,generator=g)
def rmask(*s): return (torch.rand(*s,generator=g)>0.35).float()
# replicate harness sequence up to whiten... instead just search for failing case
for i in range(200):
    v=rn(7); m=rmask(7); sh=bool(torch.randint(0,2,(1,),generator=g).item())
    if m.sum()==0: 
        r=R.whiten(v,m,sh); print("empty mask r=",r); 
        print("mine", O.whiten(v,m,sh)); break
EOF
python3 /app/_dbg.py
output
(Bash completed with no output)

Reproduce whiten failure

shell
cat > /app/_dbg.py << 'EOF'
import torch, reforacle as R, objectives as O
g=torch.Generator().manual_seed(55)
def rn(*s): return torch.randn(*s,generator=g)
def rmask(*s): return (torch.rand(*s,generator=g)>0.35).float()
for i in range(500):
    v=rn(7); m=rmask(7); sh=bool(torch.randint(0,2,(1,),generator=g).item())
    r=R.whiten(v,m,sh); mine=O.whiten(v,m,sh)
    if not torch.allclose(r,mine,atol=1e-5,rtol=1e-4,equal_nan=True):
        print("FAIL sh",sh,"masksum",m.sum().item())
        print("v",v.tolist())
        print("m",m.tolist())
        print("ref ",r.tolist())
        print("mine",mine.tolist())
        break
else: print("no fail in 500")
EOF
python3 /app/_dbg.py
output
FAIL sh False masksum 1.0
v [-0.8841636180877686, 0.5962507724761963, -0.7032215595245361, -0.6375294327735901, -0.8732007741928101, 1.6124345064163208, -2.046596050262451]
m [0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0]
ref  [-2466.979248046875, 12337.1640625, -657.5587768554688, -0.6375294327735901, -2357.350830078125, 22499.00390625, -14091.3046875]
mine [-24663418.0, 123378016.0, -6569213.0, -0.6375294327735901, -23567134.0, 224996400.0, -140906672.0]

[stdout]
FAIL sh False masksum 1.0
v [-0.8841636180877686, 0.5962507724761963, -0.7032215595245361, -0.6375294327735901, -0.8732007741928101, 1.6124345064163208, -2.046596050262451]
m [0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0]
ref  [-2466.979248046875, 12337.1640625, -657.5587768554688, -0.6375294327735901, -2357.350830078125, 22499.00390625, -14091.3046875]
mine [-24663418.0, 123378016.0, -6569213.0, -0.6375294327735901, -23567134.0, 224996400.0, -140906672.0]

/app/objectives.py

− removed
    mean = _masked_mean(values, mask)
    var = _masked_mean((values - mean) ** 2, mask)
    whitened = (values - mean) / (torch.sqrt(var) + 1e-8)
    if not shift_mean:
        whitened = whitened + mean
    return whitened
+ added
    mean = _masked_mean(values, mask)
    var = _masked_mean((values - mean) ** 2, mask)
    whitened = (values - mean) / torch.sqrt(var + 1e-8)
    if not shift_mean:
        whitened = whitened + mean
    return whitened

/app/objectives.py

− removed
    mean = _masked_mean(values, mask)
    var = _masked_mean((values - mean) ** 2, mask)
    whitened = (values - mean) / (torch.sqrt(var) + 1e-8)
    whitened = whitened * mask
+ added
    mean = _masked_mean(values, mask)
    var = _masked_mean((values - mean) ** 2, mask)
    whitened = (values - mean) / torch.sqrt(var + 1e-8)
    whitened = whitened * mask

/app/objectives.py

− removed
    adv_w = (adv - mean) / (torch.sqrt(var) + 1e-8)
+ added
    adv_w = (adv - mean) / torch.sqrt(var + 1e-8)

Rerun harness after fixes

shell
sed -i 's/^chk("entropy",logits_labels_mask)/chk("entropy",lambda:(lambda x:(x[0],x[2]))(logits_labels_mask()))/' _harness.py
sed -i 's/^chk("normalized_entropy",logits_labels_mask)/chk("normalized_entropy",lambda:(lambda x:(x[0],x[2]))(logits_labels_mask()))/' _harness.py
python3 _harness.py 2>&1 | tail -6
output
PASS first_nonzero_index 
PASS cumulative_max 
PASS bucketize_reward 

== SUMMARY ==
49/49 passing

[stdout]
PASS first_nonzero_index 
PASS cumulative_max 
PASS bucketize_reward 

== SUMMARY ==
49/49 passing

Edge-case stress test

shell
cat > /app/_edge.py << 'EOF'
import torch, reforacle as R, objectives as O
fail=[]
def eq(name,r,m,composite=False):
    atol=1e-5 if composite else 1e-6; rtol=1e-4 if composite else 1e-5
    r=r if torch.is_tensor(r) else torch.tensor(r); m=m if torch.is_tensor(m) else torch.tensor(m)
    if r.shape!=m.shape or not torch.allclose(r.float(),m.float(),atol=atol,rtol=rtol,equal_nan=True):
        fail.append((name,r,m))
        print("FAIL",name,"\n ref",r,"\n mine",m)

def C(name,*a,comp=False):
    try: eq(name,getattr(R,name)(*a),getattr(O,name)(*a),comp)
    except Exception as e: fail.append((name,e)); print("EXC",name,e)

# all-zero mask
z=torch.zeros(5); C("masked_mean",torch.randn(5),z,None); C("masked_sum",torch.randn(5),z,None)
C("masked_max",torch.randn(5),z)  # all masked -> -inf
C("whiten",torch.randn(5),z,True); C("masked_whiten",torch.randn(5),z,False)
C("advantage_mean_std",torch.randn(5),z)
# all-ones mask
o=torch.ones(4,3); C("masked_mean",torch.randn(4,3),o,1)
# constant input normalize/whiten
C("normalize",torch.full((5,),3.0),1e-8)
C("whiten",torch.full((5,),2.0),torch.ones(5),True)
# grpo/rloo group sizes
r=torch.randn(20)
for gs in [1,2,4,5,10,20]:
    C("grpo_advantages",r,gs,True); C("grpo_advantages",r,gs,False)
    C("rloo_advantages",r,gs); C("group_mean_baseline",r,gs)
# top_p boundaries
p=torch.softmax(torch.randn(3,6),-1)
for pp in [0.0,0.01,0.5,0.99,1.0]: C("top_p_mask",p,pp)
# top_k
lg=torch.randn(4,7)
for k in [1,2,7]: C("top_k_mask",lg,k)
C("top_k_mask",torch.tensor([[1.,1.,1.,1.]]),2)
# quantile edges
x=torch.randn(9)
for q in [0.0,0.25,0.5,0.75,1.0]: C("quantile_lower",x,q)
C("median_reward",torch.randn(6)); C("median_reward",torch.randn(7))
# cross_entropy with ignores
lg=torch.randn(6,5); lb=torch.randint(0,5,(6,)); lb[::2]=-100
C("cross_entropy",lg,lb,-100)
C("cross_entropy",torch.randn(3,4),torch.tensor([-100,-100,-100]),-100)
C("smoothed_nll",torch.randn(5,6),torch.randint(0,6,(5,)),0.0,comp=True)
C("smoothed_nll",torch.randn(5,6),torch.randint(0,6,(5,)),0.3,comp=True)
# bucketize exact boundaries + right
C("bucketize_reward",torch.tensor([-1.,0.,0.5,1.,2.,3.,5.]),torch.tensor([0.,1.,2.,3.]))
# first_nonzero all zero / all nonzero
C("first_nonzero_index",torch.tensor([[0.,0.,0.],[1.,1.,1.],[0.,0.,2.]]))
# cumulative_max 2d and negatives
C("cumulative_max",torch.tensor([[-1.,5.],[3.,-2.],[0.,9.]]))
C("cumulative_max",torch.randn(8))
# mode ties/unique
C("mode_label",torch.tensor([3,3,1,1,2])); C("mode_label",torch.tensor([5])); C("mode_label",torch.tensor([1,2,3,4]))
# pad mask
C("pad_mask_from_lengths",torch.tensor([0,1,5,3]),5)
# importance_ratio None + clip
C("importance_ratio",torch.randn(6),torch.randn(6),None)
C("importance_ratio",torch.randn(6),torch.randn(6),0.1)
# kl estimators
for est in ["k1","k2","k3"]: C("kl_penalty",torch.randn(7),torch.randn(7),est)
# sequence_logprob length norm with zero-length seq
lg=torch.randn(2,3,4); lb=torch.randint(0,4,(2,3)); mk=torch.tensor([[0.,0.,0.],[1.,1.,0.]])
C("sequence_logprob",lg,lb,mk,True); C("sequence_logprob",lg,lb,mk,False)
C("entropy",lg,mk); C("normalized_entropy",lg,mk); C("selective_logprobs",lg,lb,mk)
# gae/discounted 2d
C("discounted_returns",torch.randn(4,3),0.9)
C("value_loss",torch.randn(6),torch.randn(6),torch.randn(6),0.0,comp=True)
C("huber_value_loss",torch.randn(6),torch.randn(6),0.5,comp=True)
C("clip_fraction",torch.randn(10)*0.5,torch.randn(10)*0.5,0.2)
print("\nDONE fails:",len(fail))
EOF
python3 /app/_edge.py
output
FAIL masked_max 
 ref tensor(0.) 
 mine tensor(-inf)

DONE fails: 1

[stdout]
FAIL masked_max 
 ref tensor(0.) 
 mine tensor(-inf)

DONE fails: 1

/app/objectives.py

− removed
    filled = values.masked_fill(mask == 0, float("-inf"))
    return filled.max()
+ added
    filled = values.masked_fill(mask == 0, float("-inf"))
    res = filled.max()
    return torch.where(torch.isinf(res), torch.zeros_like(res), res)

Verify fix and multi-seed harness

shell
python3 _edge.py 2>&1 | tail -3 && echo "--- multi-seed main harness ---" && for s in 7 42 99 2024; do sed "s/manual_seed(1234)/manual_seed($s)/" _harness.py > _h_$s.py; python3 _h_$s.py 2>&1 | tail -1; rm _h_$s.py; done
output
DONE fails: 0
--- multi-seed main harness ---
49/49 passing
49/49 passing
49/49 passing
49/49 passing

[stdout]
DONE fails: 0
--- multi-seed main harness ---
49/49 passing
49/49 passing
49/49 passing
49/49 passing

/app/objectives.py

− removed
    adv = grpo_advantages(rewards, group_size, scale_by_std)
    adv = adv.unsqueeze(-1).expand_as(labels.float() if labels.dtype.is_floating_point
                                      else labels)
    adv = adv.to(logits.dtype)
    logp = _token_logprobs(logits, labels)
+ added
    adv = grpo_advantages(rewards, group_size, scale_by_std).unsqueeze(-1)
    logp = _token_logprobs(logits, labels)

Final verification and standalone check

shell
python3 _harness.py 2>&1 | tail -2 && python3 _edge.py 2>&1 | tail -1
# confirm standalone (no reforacle dependency) by importing in isolation
python3 -c "import ast; ast.parse(open('/app/objectives.py').read()); print('syntax OK')"
grep -c "reforacle" /app/objectives.py && echo "references present" || echo "no reforacle references (good)"
output
== SUMMARY ==
49/49 passing
DONE fails: 0
syntax OK
0
no reforacle references (good)

[stdout]
== SUMMARY ==
49/49 passing
DONE fails: 0
syntax OK
0
no reforacle references (good)

Remove temporary probe files

shell
rm -f /app/_p*.py /app/_dbg.py /app/_harness.py /app/_edge.py; ls /app
output
__pycache__
objectives.py
oracle
probe.py
reforacle.cpython-312-aarch64-linux-gnu.so
timer.sh

[stdout]
__pycache__
objectives.py
oracle
probe.py
reforacle.cpython-312-aarch64-linux-gnu.so
timer.sh
I've reimplemented all 49 functions in `/app/objectives.py`, recovering each definition by probing `reforacle` directly. Every function agrees with the reference on my test suite. ## Summary I reverse-engineered each function's exact numerical convention by calling `reforacle` and comparing, then validated with a comparison harness (49/49 passing across 5 random seeds + a dedicated edge-case suite covering empty masks, ties, degenerate variance, group-size variations, quantile/bucketize boundaries, and ignore-index handling). **Key conventions I had to discover (not assume):** - `masked_mean`/`masked_sum`/`masked_max` return **0** on all-masked input. - `whiten`/`masked_whiten` use **masked biased variance** with denominator `sqrt(var + 1e-8)` (inside the sqrt , this was a real bug in my first pass, exposed only by single-element masks). `masked_whiten` zeros masked positions *before* re-adding the mean when `shift_mean=False`. - `normalize` uses **unbiased** std with `std + eps`; `grpo_advantages` uses **biased** group std with `eps=1e-6`. - `symmetric_kl` = `cosh(d)−1` (mean of forward+reverse k3); `reverse_kl` = k3 estimator. - `cumulative_max` and `discounted_returns` operate along **dim 0** with an *exclusive* prefix (out[0]=x[0]). - `top_p_mask` keeps the token that crosses the threshold (`cumsum>=p`, shift-right) then renormalizes; `top_k_mask`/`quantile_lower` use threshold/`interpolation='lower'` semantics; `median_reward` uses linear-interpolation quantile; `mode_label` returns the **largest** among tied modes. - `bucketize_reward` uses `right=True`. - The composite objectives were the hardest , notably `reverse_kl_objective`, which has an unusual reduction: the advantage term is a per-column dot-product `-(Σ_i m·logp)·(Σ_i adv)/Σmask` and the KL term is `beta·B·masked_mean(kl)`. I recovered this by probing the linear functional with one-hot advantage vectors. `ppo_objective` whitens GAE advantages (biased) before the clipped PG loss; `grpo_objective`/`rloo_objective` broadcast group advantages per-token with a standard clipped surrogate. The file depends only on `torch` and contains no reference to `reforacle`, so it will run 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: 42/49  (partial_score=0.8571)
REWARD = 0   (1 iff ALL 49 functions match the sealed reference exactly)

The agent probes a sealed compiled oracle (reforacle) to recover 49 post-training objective conventions (DPO/IPO/GRPO/PPO/GAE/KL/RLOO) and reimplements them in torch.

Reproduce this trial: git checkout 2f94510 && PYTHONPATH=src python3 scripts/build_site.py , then open trial/trial_67a684402caa49b1. Re-running the agent live requires EVAL_PLATFORM_ENABLE_OAUTH_SMOKE=1 and is non-deterministic.

Trial trial_67a684402caa49b1 · verifier authoritative; classifier explanatory.