JunYoungLee's picture
Archive anonymized LoopQ calibration artifacts from compute node 1
d8a4a24 verified
Raw History Blame Contribute Delete
9.49 kB
"""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)}