File size: 8,140 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
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
"""Resumable real-model calibration phases, without storing backbone weights."""
from __future__ import annotations

import os
from pathlib import Path
import torch

from .calibrate import calibration_step
from .sharing_gap import SharingGapStatistics


class StatisticsAccumulator:
    """Dataset mean gradients and empirical Fisher with constant sample memory."""

    def __init__(self, state=None):
        self.count = 0 if state is None else state['count']
        self.sums = {} if state is None else state['sums']
        self.squares = {} if state is None else state['squares']
        self.parameters = {} if state is None else state['parameters']

    def add(self, statistics):
        if self.count and set(statistics) != set(self.sums):
            raise ValueError('SLT candidate set changed within a scan')
        for key, item in statistics.items():
            item.validate()
            gradients = item.gradients.detach().to(device='cpu', dtype=torch.float64)
            if not self.count:
                self.sums[key] = torch.zeros_like(gradients)
                self.squares[key] = torch.zeros_like(gradients[0])
                self.parameters[key] = item.parameter_count
            if (gradients.shape != self.sums[key].shape or
                    item.parameter_count != self.parameters[key]):
                raise ValueError('SLT candidate shape or parameter count changed')
            self.sums[key].add_(gradients)
            # Keep the Fisher estimator separate from the numerator gradients.
            # Recomputing it here would silently discard externally collected
            # score-function Fisher and force every caller onto the same proxy.
            self.squares[key].add_(item.fisher_diagonal.detach().to(device='cpu', dtype=torch.float64))
        self.count += 1

    def result(self):
        if not self.count:
            raise ValueError('cannot aggregate an empty SLT scan')
        return {key: SharingGapStatistics(self.sums[key] / self.count,
                    self.squares[key] / self.count, self.parameters[key])
                for key in self.sums}

    def state_dict(self):
        return dict(count=self.count, sums=self.sums, squares=self.squares,
                    parameters=self.parameters)


def atomic_save(value, path):
    path = Path(path)
    path.parent.mkdir(parents=True, exist_ok=True)
    temporary = path.with_name(path.name + '.tmp')
    with temporary.open('wb') as stream:
        torch.save(value, stream)
        stream.flush()
        os.fsync(stream.fileno())
    os.replace(temporary, path)


def component_modules(qdq, cta):
    return {'las': qdq.las, 'shared': qdq.shared_transforms,
            'selected': qdq.selected_loop_transforms, 'cta': cta}


class CalibrationJournal:
    """Checkpoint scan cursors, optimizer moments, component tensors and RNG.

    Progressive-selection callbacks replay saved completed scans on resume.
    An interrupted scan/optimization phase continues from its committed cursor.
    Optimizer moments intentionally reset between adaptation phases, matching
    the explicit phase policy of the real drivers, but never within a phase.
    """

    def __init__(self, path, *, contract, qdq, cta, mu_cache, interval=25, resume=False, stop_after_steps=None):
        if interval <= 0:
            raise ValueError('checkpoint interval must be positive')
        self.path, self.contract = Path(path), contract
        self.qdq, self.cta, self.mu_cache = qdq, cta, mu_cache
        self.stop_after_steps = stop_after_steps
        self.interval = interval
        self.state = dict(step=0, traces=[], optimized={}, scans={}, scan=None,
                          optimizer=None, optimizer_phase=None)
        if resume:
            saved = torch.load(self.path, map_location='cpu', weights_only=True)
            if saved.get('format') != 'loopq_real_calibration_v1' or saved['contract'] != contract:
                raise ValueError('resume source/data/config contract does not match')
            for key in saved['selected_groups']:
                qdq.select_group(key)
            for key, module in component_modules(qdq, cta).items():
                module.load_state_dict(saved['modules'][key])
            mu_cache.load_state_dict(saved['mu'])
            self.state = saved['progress']
            torch.set_rng_state(saved['rng_cpu'])
            if torch.cuda.is_available():
                torch.cuda.set_rng_state_all(saved['rng_cuda'])
        elif self.path.exists():
            raise ValueError('checkpoint exists; use --resume to continue it')
        else:
            self.save()

    @property
    def step(self):
        return self.state['step']

    @property
    def traces(self):
        return self.state['traces']

    def save(self):
        atomic_save(dict(format='loopq_real_calibration_v1', contract=self.contract,
            selected_groups=[key for key in sorted(self.qdq.sites)
                if self.qdq._encoded[key] in self.qdq.selected_loop_transforms],
            modules={key: module.state_dict() for key, module in
                     component_modules(self.qdq, self.cta).items()},
            mu=self.mu_cache.state_dict(), progress=self.state,
            rng_cpu=torch.get_rng_state(),
            rng_cuda=torch.cuda.get_rng_state_all() if torch.cuda.is_available() else []), self.path)

    def optimize(self, phase, count, optimizer_factory, loss_for):
        done = self.state['optimized'].get(phase, 0)
        if done >= count:
            return
        allowlist, optimizer = optimizer_factory()
        if self.state['optimizer_phase'] == phase:
            optimizer.load_state_dict(self.state['optimizer'])
        for index in range(done, count):
            loss = loss_for(self.step, self.mu_cache, False)[0]
            trace = calibration_step(loss=loss, optimizer=optimizer, allowlist=allowlist)
            trace.update(step=self.step, phase=phase)
            self.traces.append(trace)
            self.state['step'] += 1
            self.state['optimized'][phase] = index + 1
            self.state['optimizer_phase'] = phase
            self.state['optimizer'] = optimizer.state_dict()
            print(f'loopq optimize {phase} {index + 1}/{count} step={self.step} loss={trace["total"]:.6g}', flush=True)
            if (index + 1) % self.interval == 0 or index + 1 == count:
                self.save()
            if self.stop_after_steps is not None and self.step >= self.stop_after_steps:
                self.save()
                raise SystemExit("Requested calibration step boundary reached; checkpoint saved. Resume with --resume.")

    def scan(self, round_index, sample_count, loss_for):
        key = str(round_index)
        if key in self.state['scans']:
            return {name: SharingGapStatistics.from_dict(item)
                    for name, item in self.state['scans'][key].items()}
        pending = self.state['scan']
        if pending is not None and pending['round'] != round_index:
            raise ValueError('resume SLT round does not match saved scan')
        accumulator = StatisticsAccumulator(None if pending is None else pending['statistics'])
        for index in range(accumulator.count, sample_count):
            # Scores evaluate each sample's detached trust weights at the current
            # model. A scan never advances the optimizer clock or its mu cache.
            from .objective import AdaptiveMuCache
            scan_mu = AdaptiveMuCache(self.mu_cache.update_interval)
            loss, statistics = loss_for(index, scan_mu, True)
            accumulator.add(statistics)
            del loss, statistics
            self.state['scan'] = dict(round=round_index, statistics=accumulator.state_dict())
            print(f'loopq SLT round={round_index} sample={index + 1}/{sample_count}', flush=True)
            if (index + 1) % self.interval == 0:
                self.save()
        result = accumulator.result()
        self.state['scans'][key] = {name: item.to_dict() for name, item in result.items()}
        self.state['scan'] = None
        self.save()
        return result