File size: 8,765 Bytes
9118991 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 | """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
|