Download loopq_quantization/scripts/loopq/trajectory_diagnostics.py from JunYoungLee/ut-depth-probe-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 2.49 kB
-
https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/trajectory_diagnostics.py
- Command line
-
hf download hf://JunYoungLee/ut-depth-probe-artifacts/loopq_quantization/scripts/loopq/trajectory_diagnostics.py
-
curl -L -o trajectory_diagnostics.py https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/trajectory_diagnostics.py
2.49 kB
| """Layer-output diagnostics; reductions are explicit local measurement choices.""" | |
| from contextlib import contextmanager | |
| import torch | |
| def capture_layer_outputs(layers, loop_count): | |
| states, handles = [], [] | |
| width = len(layers) | |
| if width == 0 or loop_count <= 0: | |
| raise ValueError('positive layer and loop counts required') | |
| def capture(index, output): | |
| if index != len(states) % width or len(states) >= width * loop_count: | |
| raise ValueError('unexpected recurrent layer invocation order') | |
| value = output[0] if isinstance(output, tuple) else output | |
| if not isinstance(value, torch.Tensor) or value.ndim != 3: | |
| raise ValueError('expected batch/token/hidden layer output') | |
| states.append(value.detach().to(device='cpu', dtype=torch.float32, copy=True)) | |
| try: | |
| for index, layer in enumerate(layers): | |
| handles.append(layer.register_forward_hook( | |
| lambda module, inputs, output, index=index: capture(index, output))) | |
| yield states | |
| if len(states) != width * loop_count: | |
| raise ValueError('incomplete recurrent layer trajectory') | |
| finally: | |
| for handle in handles: | |
| handle.remove() | |
| def compare_layer_outputs(reference, observed, *, layer_count, loop_count): | |
| if layer_count <= 0 or loop_count <= 0: | |
| raise ValueError('positive layer and loop counts required') | |
| if len(reference) != layer_count * loop_count or len(observed) != len(reference): | |
| raise ValueError('incomplete recurrent layer trajectory') | |
| rows = [] | |
| for index, (ref, actual) in enumerate(zip(reference, observed)): | |
| if ref.shape != actual.shape or ref.ndim != 3 or ref.numel() == 0: | |
| raise ValueError('layer output shapes differ or are empty') | |
| ref, actual = ref.double(), actual.double() | |
| if not torch.isfinite(ref).all() or not torch.isfinite(actual).all(): | |
| raise ValueError('nonfinite layer output') | |
| delta = actual - ref | |
| norm = ref.norm().item() | |
| error = delta.norm().item() | |
| rows.append(dict(loop=index // layer_count, layer=index % layer_count, | |
| trajectory_index=index, tokens=ref.shape[0] * ref.shape[1], | |
| l2_frobenius=error, mean_token_l2=delta.norm(dim=-1).mean().item(), | |
| relative_frobenius=error / norm if norm else None, | |
| reference_frobenius=norm, max_abs_error=delta.abs().max().item())) | |
| return rows | |