Download loopq_quantization/scripts/loopq/parity_worker.py from JunYoungLee/ut-depth-probe-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 9.49 kB
-
https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/parity_worker.py
- Command line
-
hf download hf://JunYoungLee/ut-depth-probe-artifacts/loopq_quantization/scripts/loopq/parity_worker.py
-
curl -L -o parity_worker.py https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/parity_worker.py
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)} | |