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