"""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, =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 : 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)}