"""Real miniature SFT/LoRA and DPO optimization; no general capability claim.""" from pathlib import Path import json import sys import copy import time sys.path.insert(0, str(Path(__file__).resolve().parents[1])) import torch from safetensors.torch import load_file, save_file from nexora.model import NexoraLM, ModelConfig from nexora.tokenizer import ByteTokenizer from nexora.adapters import inject_lora, merge_lora from nexora.posttraining import masked_sft_loss, dpo_loss def main(): torch.set_num_threads(4) torch.manual_seed(42) root = Path("artifacts/tiny") cfg = ModelConfig(**json.loads((root / "config.json").read_text())) model = NexoraLM(cfg) model.load_state_dict(load_file(str(root / "model.safetensors"))) reference = copy.deepcopy(model).eval().requires_grad_(False) replaced = inject_lora(model) opt = torch.optim.AdamW([p for p in model.parameters() if p.requires_grad], lr=.005) tok = ByteTokenizer() rows = [ ("User: What confirms a code change?\nAssistant: ", "Run tests and inspect their actual results."), ("User: A test failed. Is the repair verified?\nAssistant: ", "No. Inspect the failure, repair it, and run the test again."), ("User: What belongs in a tool receipt?\nAssistant: ", "The operation, observed output, exit status, and elapsed time."), ("User: How should memory represent a guess?\nAssistant: ", "Record its source and uncertainty instead of storing it as a fact."), ] def tensors(prompt, answer): prefix = [tok.bos_id, *tok.encode(prompt)] ids = prefix + tok.encode(answer) + [tok.eos_id] x, y = torch.tensor([ids[:-1]]), torch.tensor([ids[1:]]) mask = torch.arange(y.shape[1])[None] >= len(prefix)-1 return x, y, mask def logprob(m, prompt, answer): x, y, mask = tensors(prompt, answer) logits, _ = m(x) return (logits.log_softmax(-1).gather(-1, y[..., None]).squeeze(-1)*mask).sum(-1) history = [] start = time.perf_counter() for step in range(40): prompt, answer = rows[step % len(rows)] x, y, mask = tensors(prompt, answer) logits, _ = model(x) loss = masked_sft_loss(logits, y, mask) opt.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1) opt.step() if step % 10 == 0 or step == 39: history.append({"stage": "SFT", "step": step+1, "loss": loss.item()}) prompt = "User: A test failed. Is the repair verified?\nAssistant: " chosen, rejected = "No. Inspect the failure and test again.", "Yes. Everything passed successfully." with torch.no_grad(): rc, rr = logprob(reference, prompt, chosen), logprob(reference, prompt, rejected) for step in range(10): loss = dpo_loss(logprob(model, prompt, chosen), logprob(model, prompt, rejected), rc, rr) opt.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1) opt.step() history.append({"stage": "DPO", "step": step+1, "loss": loss.item()}) out = Path("artifacts/posttraining-experiment") out.mkdir(exist_ok=True) save_file({k: v.detach().contiguous() for k, v in model.state_dict().items() if k.endswith((".a", ".b"))}, str(out / "adapter.safetensors")) probe, _, _ = tensors(*rows[0]) with torch.no_grad(): before, _ = model(probe) merged = merge_lora(copy.deepcopy(model)) after, _ = merged(probe) torch.testing.assert_close(before, after, atol=2e-5, rtol=2e-5) report = {"status": "VALIDATED_TOY_OPTIMIZATION_ONLY", "base": "artifacts/tiny", "rank": 4, "alpha": 8, "targets": replaced, "seconds": time.perf_counter()-start, "history": history, "merge_max_abs_error": (before-after).abs().max().item(), "limitations": "Four synthetic SFT examples and one preference pair; no evidence of reasoning improvement or generalization. Adapter not enabled by default."} (out / "config.json").write_text(json.dumps(report, indent=2)) Path("reports/posttraining.json").write_text(json.dumps(report, indent=2)) print(json.dumps(report)) if __name__ == "__main__": main()