JunYoungLee's picture
Add LoopQ 4-bit quantization of Ouro-1.4B
9118991 verified
Raw History Blame Contribute Delete
2.49 kB
"""Layer-output diagnostics; reductions are explicit local measurement choices."""
from contextlib import contextmanager
import torch
@contextmanager
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