File size: 3,593 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
"""Force-kill a worker immediately after a committed checkpoint; verify exact recovery."""
from pathlib import Path
import argparse
import json
import os
import subprocess
import sys
import tempfile
import time
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--worker")
    args = parser.parse_args()
    import torch
    from safetensors.torch import load_file
    from nexora import training
    from nexora.data import prepare
    if args.worker:
        folder = Path(args.worker)
        save = training.save_checkpoint
        def checkpoint(*a, **kw):
            receipt = save(*a, **kw)
            if receipt["step"] == 3:
                (folder / "ready").write_text("committed")
                time.sleep(120)
            return receipt
        training.save_checkpoint = checkpoint
        training.train(folder / "config.json", folder / "data", folder / "interrupted")
        return
    with tempfile.TemporaryDirectory(prefix="nexora-recovery-") as d:
        folder = Path(d)
        cfg = {"model": {"hidden_size": 32, "layers": 1, "heads": 4, "kv_heads": 2, "intermediate_size": 64, "max_context": 32},
               "training": {"steps": 9, "batch_size": 2, "sequence_length": 16, "learning_rate": .001, "seed": 73, "eval_every": 3, "checkpoint_every": 3, "device": "cpu", "threads": 2}}
        (folder / "config.json").write_text(json.dumps(cfg))
        prepare([{"id": "a", "text": "Training for recovery uses a small deterministic sequence of original engineering text.", "source": "original", "license": "owner-authored", "domain": "text"},
                 {"id": "b", "text": "Separate validation helps detect changes in the sequence of restored parameter updates.", "source": "original", "license": "owner-authored", "domain": "text", "split": "validation"}], folder / "data")
        training.train(folder / "config.json", folder / "data", folder / "full")
        with (folder / "worker.log").open("w") as log:
            worker = subprocess.Popen([sys.executable, str(Path(__file__).resolve()), "--worker", d], stdout=log, stderr=subprocess.STDOUT)
            deadline = time.monotonic()+90
            try:
                while not (folder / "ready").exists():
                    if worker.poll() is not None or time.monotonic() > deadline:
                        raise RuntimeError("Worker did not produce committed checkpoint")
                    time.sleep(.1)
                worker.kill()
                code = worker.wait(timeout=10)
            finally:
                if worker.poll() is None:
                    worker.kill()
                    worker.wait(timeout=10)
        training.train(folder / "config.json", folder / "data", folder / "interrupted", resume=True)
        a = load_file(str(folder / "full/model.safetensors"))
        b = load_file(str(folder / "interrupted/model.safetensors"))
        exact = all(torch.equal(a[k], b[k]) for k in a)
        if not exact:
            raise AssertionError("Forced-kill recovery diverged")
        report = {"status": "VALIDATED", "worker_forcibly_killed": True, "worker_exit_code": code,
                  "checkpoint_step": 3, "final_step": 9, "all_parameters_bitwise_equal": exact,
                  "limitations": "Single-process CPU recovery after checkpoint commit; not distributed kill recovery or mid-write power-loss durability"}
        Path("reports/recovery.json").write_text(json.dumps(report, indent=2))
        print(json.dumps(report))


if __name__ == "__main__":
    main()