JunYoungLee's picture
Add LoopQ 4-bit quantization of Ouro-1.4B
9118991 verified
Raw History Blame Contribute Delete
4.31 kB
"""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")