Download loopq_quantization/scripts/loopq/journal.py from JunYoungLee/ut-depth-probe-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 8.14 kB
-
https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/journal.py
- Command line
-
hf download hf://JunYoungLee/ut-depth-probe-artifacts/loopq_quantization/scripts/loopq/journal.py
-
curl -L -o journal.py https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/journal.py
8.14 kB
| """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() | |
| def step(self): | |
| return self.state['step'] | |
| 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 | |