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