File size: 4,229 Bytes
12496fc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
"""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()