Download scripts/e2e.py from Miruamel/algo-v1-codepoint: direct link, hf CLI and curl.
- Browser
- Download file 5.17 kB
-
https://huggingface.co/Miruamel/algo-v1-codepoint/resolve/main/scripts/e2e.py
- Command line
-
hf download hf://Miruamel/algo-v1-codepoint/scripts/e2e.py
-
curl -L -o e2e.py https://huggingface.co/Miruamel/algo-v1-codepoint/resolve/main/scripts/e2e.py
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() | |