JunYoungLee's picture
Add LoopQ 4-bit quantization of Ouro-1.4B
9118991 verified
Raw History Blame Contribute Delete
8.77 kB
"""Verified packed projection loading and lossless vLLM shard assembly."""
from __future__ import annotations
import hashlib
import json
from pathlib import Path
import re
import torch
from safetensors import safe_open
from adapters.ouro import PAPER_GROUPS
from .export import load_export_artifact
from .packed_weight import PackedLoopQWeight
from .paths import pinned_snapshot, REVISIONS
def _sha256(path):
digest = hashlib.sha256()
with Path(path).open('rb') as stream:
for block in iter(lambda: stream.read(1024*1024), b''):
digest.update(block)
return digest.hexdigest()
def concatenate_packed_rows(shards):
"""Preserve per-row codes/scales while assembling Q,K,V or gate,up order."""
if not shards:
raise ValueError('cannot concatenate an empty shard list')
for shard in shards:
shard.validate()
first = shards[0]
if any(s.shape[1] != first.shape[1] or s.output_dtype != first.output_dtype
or s.codes.device != first.codes.device or s.scales.device != first.scales.device for s in shards):
raise ValueError('packed shards must share input width, dtype and device')
return PackedLoopQWeight(torch.cat([s.codes for s in shards], dim=0),
torch.cat([s.scales for s in shards], dim=0),
(sum(s.shape[0] for s in shards), first.shape[1]), first.output_dtype)
def load_packed_ouro_bundle(directory, component_artifact, *, allow_diagnostic=False):
"""Return group/loop packed matrices; do not allocate dense model weights.
Keys are (canonical group, None) for shared weights and (group, loop) for
selected variants. Files and component identity are verified before use.
"""
root = Path(directory).resolve()
manifest = json.loads((root/'manifest.json').read_text())
if manifest.get('format') != 'loopq_packed_ouro_projections' or manifest.get('format_version') != 1:
raise ValueError('unsupported packed Ouro manifest')
diagnostic = manifest.get('diagnostic')
if type(diagnostic) is not bool or manifest.get('status') != ('diagnostic_exported' if diagnostic else 'exported'):
raise ValueError('packed export is incomplete or mislabeled')
if diagnostic and not allow_diagnostic:
raise ValueError('diagnostic packed bundle requires allow_diagnostic=True')
if _sha256(component_artifact) != manifest.get('component_sha256'):
raise ValueError('packed bundle component artifact hash mismatch')
component = load_export_artifact(component_artifact)
if component['model'] != manifest.get('model') or component['model'].get('revision') != REVISIONS['ouro'][1]:
raise ValueError('packed bundle model mismatch')
calibration = component.get('calibration', {})
if calibration.get('completed') is not True or (not diagnostic and calibration.get('paper_calibration') is not True):
raise ValueError('packed bundle calibration eligibility mismatch')
source = pinned_snapshot('ouro')/'model.safetensors'
if _sha256(source) != manifest.get('backbone_sha256'):
raise ValueError('packed bundle backbone hash mismatch')
groups = {f'model.layers.{layer}.{group}' for layer in range(24) for group in PAPER_GROUPS}
components = component['components']
if set(components['shared_transforms']) != groups:
raise ValueError('component must cover all Ouro groups')
variants = [(key, None) for key in sorted(groups)]
for key, loops in components['selected_loop_transforms'].items():
if key not in groups or set(map(int, loops)) != set(range(4)):
raise ValueError('invalid selected loop coverage')
variants.extend((key, loop) for loop in range(4))
expected = {}
for key, loop in variants:
prefix, group = key.rsplit('.',1)
for projection in PAPER_GROUPS[group]['hf_weights']:
expected[(key, loop, f'{prefix}.{projection}.weight')] = None
loaded = {}
with safe_open(source, framework='pt', device='cpu') as reader:
for row in manifest['weights']:
loop = row['loop']
if loop is not None and type(loop) is not int:
raise ValueError('loop index must be an integer or null')
key = (row['group'], loop, row['source_name'])
if key not in expected or key in loaded:
raise ValueError('unexpected or duplicate packed projection')
if not re.fullmatch(r'weight_\d{4}\.pt', row['path']):
raise ValueError('invalid packed file name')
path = (root/row['path']).resolve()
if path.parent != root or _sha256(path) != row['sha256']:
raise ValueError('packed file path or hash mismatch')
packed = PackedLoopQWeight.from_state_dict(torch.load(path, weights_only=True, map_location='cpu'))
if (list(packed.shape) != reader.get_slice(row['source_name']).get_shape()
or list(packed.shape) != row['shape'] or packed.output_dtype != row['dtype']
or packed.payload_bytes != row['payload_bytes']):
raise ValueError('packed projection metadata mismatch')
loaded[key] = packed
if set(loaded) != set(expected):
raise ValueError('packed bundle is missing projections')
if sum(x.payload_bytes for x in loaded.values()) != manifest['tensor_payload_bytes']:
raise ValueError('packed payload total mismatch')
assembled = {}
for key, loop in variants:
prefix, group = key.rsplit('.',1)
assembled[(key, loop)] = concatenate_packed_rows([
loaded[(key, loop, f'{prefix}.{projection}.weight')]
for projection in PAPER_GROUPS[group]['hf_weights']])
return assembled
class PackedProjectionDispatch:
"""Inference-only packed residence with a transient dense GEMM weight.
This uses the existing BF16 linear operation, not a native INT4 kernel.
Installation checks everything before releasing any dense parameters.
"""
def __init__(self, groups, parameters_by_group):
if len({id(p) for p in parameters_by_group.values()}) != len(parameters_by_group):
raise ValueError('packed groups must own distinct projection parameters')
if {key for key,loop in groups if loop is None} != set(parameters_by_group):
raise ValueError('packed shared group coverage mismatch')
prepared = {}
for (key,loop),packed in groups.items():
if key not in parameters_by_group or (loop is not None and (type(loop) is not int or not 0 <= loop < 4)):
raise ValueError('invalid packed group or loop')
parameter = parameters_by_group[key]
if parameter.requires_grad:
raise ValueError('packed dispatch requires frozen parameters')
packed.validate()
if tuple(parameter.shape) != packed.shape or str(parameter.dtype) != packed.output_dtype:
raise ValueError('packed dispatch shape/dtype mismatch (requires unsharded compatible weights)')
prepared[(key,loop)] = PackedLoopQWeight(packed.codes.to(parameter.device),
packed.scales.to(parameter.device), packed.shape, packed.output_dtype)
for key in parameters_by_group:
loops = {loop for group,loop in prepared if group == key and loop is not None}
if loops and loops != set(range(4)):
raise ValueError('selected packed group must cover all four loops')
self.parameter_ids = {key:id(p) for key,p in parameters_by_group.items()}
self.groups = prepared
self.report = {'packed_payload_bytes':sum(p.payload_bytes for p in prepared.values()),
'shared_dense_bytes_released':sum(p.numel()*p.element_size() for p in parameters_by_group.values()),
'groups':len(prepared), 'native_int4_gemm':False}
for parameter in parameters_by_group.values():
parameter.data = parameter.data.new_empty(0)
def __call__(self, projection, key, loop, value):
if torch.is_grad_enabled():
raise RuntimeError('packed projection dispatch is inference-only')
if type(loop) is not int or not 0 <= loop < 4:
raise ValueError('invalid packed recurrence index')
if id(projection.weight) != self.parameter_ids[key]:
raise ValueError('projection parameter does not match packed group')
packed = self.groups.get((key,loop), self.groups[(key,None)])
original = projection.weight.data
projection.weight.data = packed._dequantize_validated(device=value.device)
try:
return projection(value)
finally:
projection.weight.data = original