Miruamel's picture
Publish algo v1 (code-point tokenizer): 38.8M params, 16.4M tokens, loss 2.10
05f5056 verified
Raw History Blame Contribute Delete
5.17 kB
from __future__ import annotations
import json
import sys
import tempfile
from pathlib import Path
from typing import Any, Mapping, Sequence
import torch
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from algo.config import ModelConfig, SafetyConfig
from algo.data import SimpleTokenizer, synthetic_examples
from algo.eval import EvalScores, VerifierRepairLoop, evaluate_model, strict_gate
from algo.model import MiniTransformer
from algo.rsi import accept_experiment, run_rsi_cycle
from algo.tools import execute_tool_call
from algo.train import load_checkpoint, save_checkpoint
class _OracleTokenizer:
def __init__(self) -> None:
self._tokenizer = SimpleTokenizer(vocab_size=256)
self.eos_id = self._tokenizer.eos_id
def encode(self, text: str, add_eos: bool = True) -> list[int]:
return self._tokenizer.encode(text, add_eos=add_eos)
def decode(self, ids: Sequence[int]) -> str:
return self._tokenizer.decode(ids)
class _OracleModel:
training = False
def __init__(self, tokenizer: _OracleTokenizer, answers: Mapping[str, str]) -> None:
self.tokenizer = tokenizer
self.answers = answers
def to(self, device: str) -> "_OracleModel":
return self
def eval(self) -> "_OracleModel":
self.training = False
return self
def __call__(self, input_ids: torch.Tensor) -> torch.Tensor:
return torch.zeros(input_ids.shape[0], input_ids.shape[1], 256)
def generate(
self,
input_ids: torch.Tensor,
max_new_tokens: int = 64,
eos_id: int | None = None,
**_kwargs: Any,
) -> torch.Tensor:
prompt = self.tokenizer.decode(input_ids[0].tolist())
answer = self.answers[prompt]
answer_ids = self.tokenizer.encode(answer, add_eos=False)[:max_new_tokens]
suffix = torch.tensor([answer_ids], dtype=input_ids.dtype, device=input_ids.device)
return torch.cat((input_ids, suffix), dim=-1)
def main() -> None:
try:
import torch
except ImportError as exc:
raise RuntimeError("torch is required for e2e; install project dependencies") from exc
cfg = ModelConfig()
assert cfg.parameter_count_tied == 38_806_016
assert cfg.parameter_count_untied == 47_194_624
tokenizer = SimpleTokenizer()
examples = synthetic_examples(count=6)
assert examples
tiny = ModelConfig(num_layers=1, d_model=32, num_query_heads=4, num_kv_heads=2, ffn_dim=64, vocab_size=128, context_length=16)
model = MiniTransformer(tiny)
input_ids = torch.tensor([tokenizer.encode("hello")])
with torch.no_grad():
logits = model(input_ids)
assert tuple(logits.shape) == (1, input_ids.shape[-1], tiny.vocab_size)
generated = model.generate(input_ids, max_new_tokens=4, eos_id=tokenizer.eos_id)
assert generated.shape[-1] >= input_ids.shape[-1]
memory_bytes = model.estimate_inference_memory_bytes(batch_size=1, seq_len=tiny.context_length)
assert memory_bytes <= 500 * 1024 * 1024
tool_result = execute_tool_call({"kind": "python", "code": "print(2+3)"})
assert tool_result.ok and "5" in tool_result.output
capped = execute_tool_call({"kind": "python", "code": "print('x' * 1000)"}, safety=SafetyConfig(max_output_bytes=10))
assert len(capped.output.encode("utf-8")) == 10
eval_examples = [
{"domain": "math", "text": "1+1", "answer": "2"},
{"domain": "coding", "text": "print(1)", "answer": "print(1)"},
{"domain": "tool", "text": "tool", "tool_call": {"kind": "python", "code": "print(1)"}},
]
eval_tokenizer = _OracleTokenizer()
eval_model = _OracleModel(eval_tokenizer, {"1+1": "2", "print(1)": "print(1)"})
scores = evaluate_model(
eval_model,
eval_tokenizer,
eval_examples,
tool_executor=lambda kind, payload: execute_tool_call({"kind": kind, "payload": payload}),
)
assert strict_gate(scores)
loop = VerifierRepairLoop(max_repairs=1)
attempts: list[str] = []
repaired, ok = loop.run("bad", lambda answer: (attempts.append(answer) or len(attempts) == 2), lambda answer, _feedback: answer + " repaired")
assert ok and attempts == ["bad", "bad repaired"]
baseline = EvalScores(math=0.70, coding=0.65, tool=0.85)
candidate = EvalScores(math=0.73, coding=0.67, tool=0.87)
calls: list[str] = []
cycle = run_rsi_cycle(
baseline,
candidate,
apply_candidate=lambda: calls.append("apply"),
evaluate=lambda: (_ for _ in ()).throw(RuntimeError("evaluate failed")),
revert_candidate=lambda: calls.append("revert"),
)
assert not cycle.accepted and calls == ["apply", "revert"]
assert accept_experiment(baseline, candidate).accepted
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "state.pt"
state = {"step": 3, "tokens": 123}
save_checkpoint(path, state)
assert load_checkpoint(path) == state
print(json.dumps({"status": "ok", "parameter_count_tied": cfg.parameter_count_tied, "tiny_inference_memory_bytes": memory_bytes}, sort_keys=True))
if __name__ == "__main__":
main()