File size: 9,493 Bytes
9118991
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d8a4a24
9118991
 
 
 
 
 
 
 
 
 
d8a4a24
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
"""Named vLLM diagnostic RPCs; no arbitrary-callable serialization required."""
import torch


class ParityWorkerExtension:
    def loopq_install_packed(self, directory, allow_diagnostic=False, compare_weights=False):
        """Install verified packed weights after loading, for eager single-GPU diagnostics."""
        if not self.model_runner.model_config.enforce_eager:
            raise ValueError('packed diagnostic requires eager execution')
        if self.parallel_config.tensor_parallel_size != 1 or self.parallel_config.pipeline_parallel_size != 1:
            raise ValueError('packed diagnostic requires an unsharded single GPU')
        runtime = self.model_runner.model
        bridge = runtime.model.loopq_bridge
        if bridge.artifact_path is None:
            raise ValueError('packed diagnostic requires a calibrated component artifact')
        from loopq.packed_bundle import load_packed_ouro_bundle
        groups = load_packed_ouro_bundle(directory, bridge.artifact_path,
                                        allow_diagnostic=allow_diagnostic)
        from adapters.ouro import PAPER_GROUPS
        parameters = dict(runtime.named_parameters())
        comparisons = []
        if compare_weights:
            for (key, loop), packed in groups.items():
                prefix, group = key.rsplit('.', 1)
                consumer, = PAPER_GROUPS[group]['vllm_consumers']
                dense = (parameters[f'{prefix}.{consumer}.weight'] if loop is None
                         else bridge.weight_variants[(key, loop)])
                restored = packed.dequantize(device=dense.device)
                if restored.shape != dense.shape or restored.dtype != dense.dtype:
                    raise ValueError('packed/dense diagnostic shape or dtype mismatch')
                difference = restored.float() - dense.detach().float()
                comparisons.append(dict(group=key, loop=loop, elements=dense.numel(),
                    unequal_elements=int(torch.count_nonzero(difference).item()),
                    max_abs_error=float(difference.abs().max().item())))
                del restored, difference, dense
        before = torch.cuda.memory_allocated()
        report = bridge.install_packed_groups(groups, parameters)
        report['dense_weight_comparison'] = comparisons
        torch.cuda.synchronize()
        return report | {'cuda_allocated_before': before,
                         'cuda_allocated_after': torch.cuda.memory_allocated()}

    def loopq_install_parity(self, model_name, <redacted-hf-token>=False, initial_state_path=None, attention_replay_path=None):
        runtime = self.model_runner.model
        replay = None if attention_replay_path is None else torch.load(
            attention_replay_path, weights_only=True, map_location='cpu')
        self._loopq_replay_seen = set()
        self._loopq_replay_expected = set() if replay is None else {
            (tokens, layer, loop) for tokens, layers in replay.items()
            for layer, values in layers.items() for loop in range(values.shape[0])}
        if replay is not None and model_name != 'ouro':
            raise ValueError('attention replay is Ouro-only')
        changed = []
        if <redacted-hf-token>:
            for name, module in runtime.named_modules():
                if type(module).__name__ == 'RMSNorm':
                    module.forward = module.forward_native
                    changed.append(name)
                elif type(module).__name__ == 'SiluAndMul':
                    def separate_silu(value):
                        gate, up = value.chunk(2, dim=-1)
                        return torch.nn.functional.silu(gate) * up
                    module.forward = separate_silu
                    changed.append(name)
                for flag in ['defer_interlayer_residual', 'defer_loop_boundary_residual',
                             'inplace_fused_residual_norm', 'outplace_fused_residual_norm']:
                    if hasattr(module, flag):
                        setattr(module, flag, False)
        if initial_state_path is not None:
            initial = torch.load(initial_state_path, weights_only=True, map_location='cpu')
            def provider(input_ids, positions, input_embeds):
                state = initial.to(input_embeds)
                return state.index_select(0, positions.to(device=state.device, dtype=torch.long))
            runtime.model.set_runtime_initial_state_provider(provider)
        self._loopq_parity_states = []
        self._loopq_parity_sites = {}
        self._loopq_parity_site_handles = []
        self._loopq_parity_handle = None
        if model_name == 'ouro':
            modules = {f'layer{index}_input': runtime.model.layers[index].input_layernorm
                       for index in [0, 1, 2, 4, 12, 23]}
            modules.update(qkv=runtime.model.layers[0].self_attn.qkv_proj,
                           gate_up=runtime.model.layers[0].mlp.gate_up_proj,
                           down=runtime.model.layers[0].mlp.down_proj,
                           layer0_o_projection=runtime.model.layers[0].self_attn.o_proj,
                           layer0_post_attn=runtime.model.layers[0].post_attention_layernorm)
            for index in [1, 2]:
                layer = runtime.model.layers[index]
                modules.update({f'layer{index}_qkv': layer.self_attn.qkv_proj,
                    f'layer{index}_gate_up': layer.mlp.gate_up_proj,
                    f'layer{index}_down': layer.mlp.down_proj,
                    f'layer{index}_post_attn': layer.post_attention_layernorm})
            for key, module in modules.items():
                self._loopq_parity_sites[key] = []
                def capture_site(module, inputs, output, key=key):
                    value = output[0] if isinstance(output, tuple) else output
                    self._loopq_parity_sites[key].append({'tokens': value.shape[0],
                        'value': value.detach().float().cpu().tolist()})
                self._loopq_parity_site_handles.append(module.register_forward_hook(capture_site))
            bridge = runtime.model.loopq_bridge
            self._loopq_parity_bridge = bridge
            self._loopq_original_prepare_activation = bridge.prepare_activation
            attention_layers = list(range(24)) if replay is not None else [0, 1, 2]
            for index in attention_layers:
                self._loopq_parity_sites[f'layer{index}_attention_output'] = []
            if bridge.enabled:
                original = bridge.prepare_activation
                def capture_prepare(value, *, layer_idx, loop_idx, group):
                    if layer_idx in attention_layers and group == 'attention_output':
                        self._loopq_parity_sites[f'layer{layer_idx}_attention_output'].append({
                            'tokens': value.shape[0], 'value': value.detach().float().cpu().tolist()})
                    if replay is not None and group == 'attention_output':
                        replay_key = (value.shape[0], layer_idx, loop_idx)
                        if replay_key in self._loopq_replay_seen:
                            raise ValueError('duplicate attention replay boundary')
                        self._loopq_replay_seen.add(replay_key)
                        reference = replay[value.shape[0]][layer_idx][loop_idx]
                        if reference.shape != value.shape:
                            raise ValueError('replayed attention shape mismatch')
                        value = reference.to(value)
                    return original(value, layer_idx=layer_idx, loop_idx=loop_idx, group=group)
                bridge.prepare_activation = capture_prepare
            else:
                for index in [0, 1, 2]:
                    key = f'layer{index}_attention_output'
                    def capture_attention(module, inputs, key=key):
                        value = inputs[0]
                        self._loopq_parity_sites[key].append({'tokens': value.shape[0],
                            'value': value.detach().float().cpu().tolist()})
                    self._loopq_parity_site_handles.append(runtime.model.layers[index].self_attn.o_proj.register_forward_pre_hook(capture_attention))
            def capture(module, inputs, output):
                value = output[0] if isinstance(output, tuple) else output
                self._loopq_parity_states.append(value.detach().float().cpu())
            self._loopq_parity_handle = runtime.model.norm.register_forward_hook(capture)
        return {'elementwise_changes': changed, 'initial_state_supplied': initial_state_path is not None}

    def loopq_collect_parity(self):
        if hasattr(self, "_loopq_parity_bridge"):
            self._loopq_parity_bridge.prepare_activation = self._loopq_original_prepare_activation
        if self._loopq_parity_handle is not None:
            self._loopq_parity_handle.remove()
            self._loopq_parity_handle = None
        for handle in self._loopq_parity_site_handles:
            handle.remove()
        states = [state.tolist() for state in self._loopq_parity_states]
        self._loopq_parity_states.clear()
        if self._loopq_replay_seen != self._loopq_replay_expected:
            raise ValueError('attention replay did not visit every recorded boundary exactly once')
        return {'boundaries': states, 'sites': self._loopq_parity_sites,
                'replayed_boundaries': len(self._loopq_replay_seen)}