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
|