NEXORA / scripts /posttrain_experiment.py
devildasdf's picture
Release validated NEXORA research prototype, tiny weights and evidence
12496fc verified
Raw History Blame Contribute Delete
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()