File size: 4,306 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
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
"""Exact tensor accounting; never equate payload with measured loaded memory."""
import math
from collections.abc import Mapping

import torch


def packed_w4_bytes(rows: int, width: int) -> int:
    if rows <= 0 or width <= 0:
        raise ValueError("weight dimensions must be positive")
    return rows * (math.ceil(width / 2) + 4 * math.ceil(width / 32))


def kronecker_counts(left: int, right: int) -> dict:
    if min(left, right) <= 0:
        raise ValueError("factor dimensions must be positive")
    return dict(effective_factors=left * left + right * right,
                svd_trainable=2 * left * left + left + 2 * right * right + right)


def ouro_analytic_accounting(*, layers=24, hidden=2048, intermediate=5632,
                             vocab=49152, heads=16, kv_heads=16,
                             hidden_factors=(32, 64), intermediate_factors=(64, 88),
                             loops=4, selected_groups=4, rank=8) -> dict:
    """Known Ouro layout, untied embeddings, bias-free projections, gate h+1.

    Bounds allow every selected group to be either the smallest or largest
    candidate. They include all T copies, NOT an extra-copy delta. Shared and
    loop-dependent counts remain separate from the BF16 backbone tensor count.
    """
    if hidden % heads or hidden_factors[0] * hidden_factors[1] != hidden or intermediate_factors[0] * intermediate_factors[1] != intermediate:
        raise ValueError("invalid head/factor configuration")
    kv = hidden // heads * kv_heads
    shapes = [(hidden, hidden), (kv, hidden), (kv, hidden), (hidden, hidden),
              (intermediate, hidden), (intermediate, hidden), (hidden, intermediate)]
    weight_elements = layers * sum(r * c for r, c in shapes)
    scales = layers * sum(r * math.ceil(c / 32) for r, c in shapes)
    # HF stores four norms per layer, including *_layernorm_2, even where
    # forward only consumes two. Stored payload must count the unused tensors.
    non_projection = 2 * vocab * hidden + (4 * layers + 1) * hidden + hidden + 1
    base = weight_elements + non_projection
    h, m = kronecker_counts(*hidden_factors), kronecker_counts(*intermediate_factors)
    shared_uv = 2 * hidden * rank
    dependent_cta = (loops - 1) * (2 * hidden + rank)
    las = layers * 4 * loops
    bounds = {name: [selected_groups * loops * min(h[name], m[name]) + las + dependent_cta,
                     selected_groups * loops * max(h[name], m[name]) + las + dependent_cta]
              for name in h}
    shared = {name: layers * (3 * h[name] + m[name]) for name in h}
    packed_base = layers * sum(packed_w4_bytes(*shape) for shape in shapes)
    return dict(projection_elements=weight_elements, group32_weight_scales=scales,
                backbone_stored_elements=base, bf16_payload_bytes=base * 2,
                bf16_payload_GiB=base * 2 / 2 ** 30,
                packed_projection_payload_bytes=packed_base,
                non_projection_bf16_payload_bytes=non_projection * 2,
                shared_transform_counts=shared, shared_cta_uv=shared_uv,
                loop_dependent_las=las, loop_dependent_cta=dependent_cta,
                loop_dependent_bounds=bounds,
                assumptions="fixed local Kronecker layout, all loop copies, no optional diagonals; not measured loaded memory")


def tensor_tree_inventory(tree) -> dict:
    """Count tensor objects once; storage bytes deduplicate aliased views."""
    objects, storages = {}, {}

    def visit(value):
        if isinstance(value, torch.Tensor):
            objects[id(value)] = value
            storage = value.untyped_storage()
            storages[(str(value.device), storage.data_ptr(), storage.nbytes())] = storage.nbytes()
        elif isinstance(value, Mapping):
            for child in value.values():
                visit(child)
        elif isinstance(value, (list, tuple)):
            for child in value:
                visit(child)

    visit(tree)
    return dict(tensor_objects=len(objects), tensor_elements=sum(t.numel() for t in objects.values()),
                logical_tensor_bytes=sum(t.numel() * t.element_size() for t in objects.values()),
                unique_storage_bytes=sum(storages.values()),
                excludes="allocator overhead, runtime workspaces, KV cache, Python containers")