JunYoungLee's picture
Add LoopQ 4-bit quantization of Ouro-1.4B
9118991 verified
Raw History Blame Contribute Delete
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()
@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