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)}
|