Download scripts/posttrain_experiment.py from devildasdf/NEXORA: direct link, hf CLI and curl.
- Browser
- Download file 4.23 kB
-
https://huggingface.co/devildasdf/NEXORA/resolve/main/scripts/posttrain_experiment.py
- Command line
-
hf download hf://devildasdf/NEXORA/scripts/posttrain_experiment.py
-
curl -L -o posttrain_experiment.py https://huggingface.co/devildasdf/NEXORA/resolve/main/scripts/posttrain_experiment.py
4.23 kB
| """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() | |