File size: 2,488 Bytes
9118991
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
"""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