File size: 2,079 Bytes
9118991
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Capture the exact source and input identity of a calibration run."""
import hashlib
import json
from pathlib import Path
import torch


def capture_sources(output: Path, args) -> dict:
    source_root = Path(__file__).resolve().parents[1]
    destination = output.with_suffix(".sources")
    resuming = getattr(args, "resume", False)
    if not resuming:
        destination.mkdir(parents=True, exist_ok=False)
    hashes = {}
    for path in sorted(source_root.rglob("*")):
        if path.suffix not in {".py", ".yaml"} or "__pycache__" in path.parts:
            continue
        relative = path.relative_to(source_root)
        data = path.read_bytes()
        hashes[str(relative)] = hashlib.sha256(data).hexdigest()
        target = destination / relative
        if not resuming:
            target.parent.mkdir(parents=True, exist_ok=True)
            target.write_bytes(data)
    metadata = {
        "source_sha256": hashes, "source_snapshot": str(destination.resolve()),
        "arguments": {k: str(v) if isinstance(v, Path) else v for k,v in vars(args).items()},
        "torch": str(torch.__version__), "cuda": torch.version.cuda,
        "physical_gpu": args.gpu, "gpu_count": 1,
    }
    metadata = json.loads(json.dumps(metadata))
    if resuming:
        previous = json.loads((destination / "manifest.json").read_text())
        ignored = {"resume", "gpu", "stop_after_steps"}
        requested = {key: value for key, value in metadata["arguments"].items() if key not in ignored}
        original = {key: value for key, value in previous["arguments"].items() if key not in ignored}
        # JSON normalizes factor tuples to lists.
        if json.loads(json.dumps(requested)) != original or previous["source_sha256"] != hashes:
            raise ValueError("resume requires unchanged calibration source and arguments")
        return previous
    (destination / "manifest.json").write_text(json.dumps(metadata, indent=2) + "\n")
    return metadata


def text_hashes(texts):
    return [hashlib.sha256(text.encode("utf-8")).hexdigest() for text in texts]