Download loopq_quantization/scripts/loopq/packed_bundle.py from JunYoungLee/ut-depth-probe-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 8.77 kB
-
https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/packed_bundle.py
- Command line
-
hf download hf://JunYoungLee/ut-depth-probe-artifacts/loopq_quantization/scripts/loopq/packed_bundle.py
-
curl -L -o packed_bundle.py https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/packed_bundle.py
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 | |