"""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")