"""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