Repository navigation
Expand file tree
/
Copy pathpointer.py
More file actions
189 lines (159 loc) · 8.24 KB
/
Copy pathpointer.py
File metadata and controls
189 lines (159 loc) · 8.24 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
"""Explicit-step pointer programs and exact, one-based execution targets."""
from dataclasses import asdict, dataclass
import hashlib
import json
import random
import re
from typing import Sequence
SYMBOLS = tuple("ABCDEFGHIJKLMNOPQRSTUVWXYZ")
PAIR = re.compile(r"\(\s*([A-Z])\s*,\s*([A-Z])\s*\)")
def parse_mapping(text: str) -> dict[str, str]:
"""Parse pairs such as '(A,C) (C,D) (D, E)', rejecting ambiguous rules."""
matches = list(PAIR.finditer(text))
if not matches or PAIR.sub("", text).strip():
raise ValueError("Mapping must be whitespace-separated (A,C) pairs")
mapping = {}
for match in matches:
source, target = match.groups()
if source in mapping:
raise ValueError(f"Duplicate source symbol: {source}")
mapping[source] = target
return mapping
def execute(mapping: dict[str, str], start: str, steps: int) -> list[str]:
"""Return states after transitions 1..steps, excluding the initial state.
This reference interpreter permits cycles and partial tables, provided every
requested transition exists. The dataset sampler applies stricter constraints.
"""
if type(steps) is not int or steps < 0:
raise ValueError("steps must be a nonnegative integer")
states = []
current = start
for step in range(1, steps + 1):
if current not in mapping:
raise ValueError(f"No rule for {current!r} at step {step}")
current = mapping[current]
states.append(current)
return states
def check_predictions(
mapping: dict[str, str], start: str, predictions: Sequence[str], steps: int,
) -> list[bool]:
"""Score decoded symbols at steps 1..N against the exact nominal trajectory.
Predictions may be a prefix (too few model loops), but cannot extend beyond
the requested depth. Invalid/multi-symbol answers are incorrect; no lenient
extraction of a letter from prose is performed. Leading whitespace is allowed.
"""
targets = execute(mapping, start, steps)
if isinstance(predictions, str) or len(predictions) > steps:
raise ValueError("Provide a sequence of at most 'steps' decoded predictions")
return [prediction.strip() == target for prediction, target in zip(predictions, targets)]
def render_prompt(pairs: list[list[str]], start: str, steps: int) -> str:
# Spaces before symbols prevent Qwen punctuation/letter token merges.
rules = " ".join(f"( {source}, {target})" for source, target in pairs)
return f"Rules: {rules}\nStart: {start}\nSteps: {steps}\nAnswer:"
def mapping_fingerprint(mapping: dict[str, str]) -> str:
"""Ignore rule order, start, and depth when detecting reused rule tables."""
payload = json.dumps(sorted(mapping.items()), separators=(",", ":"))
return hashlib.sha256(payload.encode()).hexdigest()
def example_seed(master_seed: int, split: str, index: int, attempt: int = 0) -> int:
payload = f"pointer-v1:{master_seed}:{split}:{index}:{attempt}"
return int.from_bytes(hashlib.sha256(payload.encode()).digest()[:8], "big")
@dataclass
class PointerExample:
schema_version: int
example_id: str
family: str
split: str
seed: int
task_depth: int
mapping: list[list[str]]
mapping_sha256: str
initial_state: str
intermediate_states: list[str]
final_state: str
prompt: str
def to_dict(self) -> dict:
return asdict(self)
def generate_example(seed: int, depth: int, split: str, index: int) -> PointerExample:
"""Sample a full table, conditioned on a non-repeating nominal path.
Every example uses all 26 symbols, keeping rule count independent of depth.
The d+1 distinct path states prevent an early/repeated final answer; remaining
edges are independent random choices and can cycle outside the nominal path.
Rule order is shuffled so the solution is not printed in execution order.
"""
if type(depth) is not int or not 1 <= depth < len(SYMBOLS):
raise ValueError("depth must be an integer between 1 and 25")
if type(seed) is not int or seed < 0:
raise ValueError("seed must be a nonnegative integer")
rng = random.Random(seed)
path = rng.sample(SYMBOLS, depth + 1)
mapping = dict(zip(path[:-1], path[1:]))
for symbol in SYMBOLS:
if symbol not in mapping:
mapping[symbol] = rng.choice(SYMBOLS)
pairs = [[source, mapping[source]] for source in SYMBOLS]
rng.shuffle(pairs)
states = execute(mapping, path[0], depth)
return PointerExample(
schema_version=1, example_id=f"{split}-{index:06d}-{seed:016x}",
family="pointer", split=split, seed=seed, task_depth=depth,
mapping=pairs, mapping_sha256=mapping_fingerprint(mapping),
initial_state=path[0], intermediate_states=states, final_state=states[-1],
prompt=render_prompt(pairs, path[0], depth),
)
def validate_example(example: PointerExample) -> None:
"""Reparse the actual prompt and independently recompute every saved target."""
if example.schema_version not in (1, 2) or example.family != "pointer":
raise ValueError("Unsupported pointer record schema/family")
if type(example.task_depth) is not int or not 1 <= example.task_depth <= (25 if example.schema_version == 1 else 256):
raise ValueError("Invalid task depth")
if any(len(pair) != 2 for pair in example.mapping):
raise ValueError("Every mapping entry must have two symbols")
mapping = dict(example.mapping)
if len(example.mapping) != len(SYMBOLS) or set(mapping) != set(SYMBOLS):
raise ValueError("Each vocabulary symbol must have exactly one rule")
if not set(mapping.values()) <= set(SYMBOLS):
raise ValueError("Mapping target outside vocabulary")
if example.mapping_sha256 != mapping_fingerprint(mapping):
raise ValueError("Mapping fingerprint mismatch")
if example.prompt != render_prompt(example.mapping, example.initial_state, example.task_depth):
raise ValueError("Prompt disagrees with structured task")
parsed = parse_mapping(example.prompt.splitlines()[0].removeprefix("Rules: "))
states = execute(parsed, example.initial_state, example.task_depth)
if states != example.intermediate_states or example.final_state != states[-1]:
raise ValueError("Intermediate/final labels disagree with reference execution")
if example.schema_version == 1 and len(set([example.initial_state, *states])) != example.task_depth + 1:
raise ValueError("Nominal path must not repeat a state")
def generate_unconditioned_example(seed: int, depth: int, split: str, index: int,
graph_mode: str = "mixture") -> PointerExample:
"""Sample the graph/start/order before using depth; permit cyclic execution.
Equal-probability mixture: random function, random permutation, full cycle.
All three have 26 rules; graph selection never consults requested horizon.
"""
if type(depth) is not int or not 1 <= depth <= 256 or type(seed) is not int or seed < 0:
raise ValueError("Need nonnegative seed and depth in 1..256")
rng = random.Random(seed)
kinds = ("random_function", "permutation", "full_cycle")
kind = rng.choice(kinds) if graph_mode == "mixture" else graph_mode
if kind not in kinds:
raise ValueError("Unknown depth-independent graph mode")
if kind == "random_function":
mapping = {s: rng.choice(SYMBOLS) for s in SYMBOLS}
else:
order = rng.sample(SYMBOLS, len(SYMBOLS))
mapping = (dict(zip(SYMBOLS, order)) if kind == "permutation" else
{s: order[(i + 1) % len(order)] for i, s in enumerate(order)})
start = rng.choice(SYMBOLS)
pairs = [[s, mapping[s]] for s in SYMBOLS]
rng.shuffle(pairs)
states = execute(mapping, start, depth)
return PointerExample(2, f"{split}-{index:06d}-{seed:016x}", "pointer", split, seed,
depth, pairs, mapping_fingerprint(mapping), start, states, states[-1],
render_prompt(pairs, start, depth))
def orbit_structure(mapping: dict[str, str], start: str) -> dict[str, int]:
"""Length before the first cycle and cycle period, independent of horizon."""
seen = {}
state = start
while state not in seen:
seen[state] = len(seen)
state = mapping[state]
return {"transient_length": seen[state], "cycle_period": len(seen) - seen[state]}