File size: 6,449 Bytes
932bc69
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
"""Full-run launcher: transactional checkpoints and durable progress metadata."""

import json
import logging
import os
import runpy
import shutil
import tempfile
import time
from collections.abc import Mapping
from pathlib import Path

import torch.distributed as dist

from ngram_checkpoint_transaction import (
    MARKER,
    latest_managed_epoch,
    publish_checkpoint,
)
from ngram_epoch_accounting import full_epoch_steps, validate_completed_state
from speculators.train.checkpointer import BaseCheckpointer
from ngram_wandb_resume import update_resumable_config
from speculators.train.logger import WandbHandler, _flatten_dict
from speculators.train.trainer import Trainer

run_root = Path(os.environ["NGRAM_RUN_ROOT"])
log_root = Path(os.environ["LOG_DIR"])
original_save = Trainer.maybe_save_checkpoint
original_setup = Trainer.setup_optimizer
original_run = Trainer.run_training
original_wandb_setup = WandbHandler._setup
original_wandb_emit = WandbHandler.emit
original_previous_epoch = BaseCheckpointer._get_previous_epoch


def previous_epoch(self):
    # Upstream ignores all symlinks. Permit only our validated numeric pointers;
    # descriptive aliases and arbitrary external symlinks remain ignored.
    return latest_managed_epoch(self.path, original_previous_epoch(self))


def wandb_setup(self):
    # Segment-local paths change at renewal. Keep the same W&B run while the
    # immutable source/config and per-segment provenance preserve all versions.
    self.init_kwargs.setdefault("allow_val_change", True)
    return original_wandb_setup(self)


def wandb_emit(self, record):
    if getattr(record, "hparams", False) and isinstance(record.msg, Mapping):
        if self._run is None:
            self._run = self._setup()
        update_resumable_config(self._run.config, _flatten_dict(record.msg))
        return
    return original_wandb_emit(self, record)


def write_json(path, value):
    temporary = path.with_suffix(".json.tmp")
    temporary.write_text(json.dumps(value, indent=2) + "\n")
    os.replace(temporary, path)


class ProgressRecorder(logging.Handler):
    def emit(self, record):
        if isinstance(record.msg, dict) and "global_step" in record.msg:
            write_json(
                run_root / "progress.json",
                {
                    "time": time.time(),
                    "segment": str(log_root),
                    "global_step": record.msg["global_step"],
                    "train": record.msg.get("train", {}),
                    "profile": record.msg.get("profile", {}),
                },
            )


def setup(self):
    original_setup(self)
    if self.rank == 0:
        for name in ("train_command.txt", "speculators.patch", "run.yaml"):
            path = self.checkpointer.path / name
            if path.is_file():
                shutil.copy2(path, log_root / name)
        if (self.config.max_steps is not None
                or self.config.num_epochs != int(os.environ["EPOCHS"])):
            raise ValueError(
                "Run must match the manifest epoch count, with no max_steps"
            )
        logging.getLogger("speculators.metrics").addHandler(ProgressRecorder())
        write_json(
            log_root / "training_start.json",
            {
                "epoch_steps": len(self.train_loader),
                "epochs": self.config.num_epochs,
                "total_steps": self.config.num_epochs * len(self.train_loader),
                "resume_global_step": self.global_step,
                "resume_local_step": self._resume_local_step,
                "current_epoch": self.current_epoch,
                "checkpoint": str(self.checkpointer.prev_path),
                "scheduler_last_epochs": [s.last_epoch for s in self.schedulers],
            },
        )


def transactional_save(self, epoch, local_step=0):
    if not isinstance(epoch, int):
        return original_save(self, epoch, local_step)
    if self.config.save_best or self.config.checkpoint_freq >= 1:
        raise ValueError("Use periodic checkpoint_freq < 1 without save_best")
    root = self.checkpointer.path
    paths = [None]
    if self.rank == 0:
        stage = Path(tempfile.mkdtemp(prefix=".pending-", dir=root))
        (stage / MARKER).touch()
        for name in ("train_command.txt", "speculators.patch", "run.yaml"):
            if (root / name).is_file():
                shutil.copy2(root / name, stage / name)
        paths[0] = str(stage)
    if self.is_distributed:
        dist.broadcast_object_list(paths, src=0)
    stage = Path(paths[0])
    self.checkpointer.path = stage
    try:
        original_save(self, epoch, local_step)
    finally:
        self.checkpointer.path = root
    if self.is_distributed:
        dist.barrier()
    if self.rank == 0:
        result = publish_checkpoint(root, stage, epoch, self.global_step)
        result.update(time=time.time(), segment=str(log_root))
        write_json(run_root / "checkpoint_commit.json", result)
        print("ATOMIC_CHECKPOINT " + json.dumps(result), flush=True)
    if self.is_distributed:
        dist.barrier()


def run(self):
    original_run(self)
    if self.rank == 0:
        final_epoch = self.config.num_epochs - 1
        checkpoint = self.checkpointer.path / str(final_epoch)
        state = json.loads(
            (checkpoint / "training_state.json").read_text()
        )
        # Packed batch counts can differ by epoch, and a resumed sampler's cache
        # contains only its remaining slice. Audit full epochs on a separate
        # sampler; never infer completion by multiplying the final loader length.
        epoch_steps = full_epoch_steps(self.train_loader, self.config.num_epochs)
        validate_completed_state(state, self.config.num_epochs, epoch_steps)
        write_json(
            run_root / "training_complete.json",
            {
                "passed": True,
                "time": time.time(),
                "state": state,
                "checkpoint": str(checkpoint),
                "epoch_steps": epoch_steps,
                "total_steps": sum(epoch_steps),
            },
        )


Trainer.setup_optimizer = setup
Trainer.maybe_save_checkpoint = transactional_save
Trainer.run_training = run
WandbHandler._setup = wandb_setup
WandbHandler.emit = wandb_emit
BaseCheckpointer._get_previous_epoch = previous_epoch
runpy.run_path(
    str(Path(__file__).resolve().parents[1] / "train.py"), run_name="__main__"
)