File size: 4,261 Bytes
13b1a91
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
"""Conservative memory estimates. Every number here is an estimate, not a measurement."""

from __future__ import annotations

from dataclasses import dataclass, asdict

from .parsing import Arch

GB = 1024**3

RAM_CLASSES = [8, 16, 18, 24, 32, 36, 48, 64, 96, 128, 192, 256, 512]
CONTEXTS = [4096, 8192, 16384, 32768, 65536, 131072, 262144]

FIT_ORDER = ["Comfortable", "Likely", "Borderline", "Unlikely"]


def bits_per_weight(bits: float | None, mode: str | None = None) -> float | None:
    """Effective storage bits including per-group scales/biases."""
    if bits is None:
        return None
    if bits >= 16:
        return 16.0
    if mode in ("mxfp4", "nvfp4", "mxfp8"):
        return bits + 0.5  # one 8-bit scale per 16-32 weights
    return bits + 0.5  # affine gs64: fp16 scale + bias per 64 weights


def weights_bytes(
    file_bytes: int | None, params: float | None, bits: float | None, mode: str | None = None
) -> tuple[int | None, str]:
    if isinstance(file_bytes, (int, float)) and file_bytes > 0:
        return int(file_bytes), "files"
    bpw = bits_per_weight(bits, mode)
    if params and bpw:
        return int(params * bpw / 8), "params"
    return None, "unknown"


def kv_cache_bytes(arch: Arch, context: int, kv_bytes_per_elem: float = 2.0) -> tuple[int | None, bool]:
    """Returns (bytes, is_upper_bound). fp16 K and V for every cached token."""
    if not arch.known:
        return None, True
    per_token_layer = 2 * arch.kv_heads * arch.head_dim * kv_bytes_per_elem
    full = arch.full_attention_layers if arch.full_attention_layers is not None else arch.layers
    total = full * context * per_token_layer
    if arch.sliding_layers and arch.sliding_window:
        total += arch.sliding_layers * min(context, arch.sliding_window) * per_token_layer
    upper = arch.full_attention_layers is None
    return int(total), upper


def usable_gpu_bytes(ram_gb: float) -> int:
    """macOS lets Metal wire roughly 2/3 of RAM on small machines and 3/4 on larger ones
    by default (raisable with `sudo sysctl iogpu.wired_limit_mb`)."""
    frac = 0.67 if ram_gb <= 36 else 0.75
    return int(ram_gb * GB * frac)


def overhead_bytes(weights: int | None) -> int:
    return int(1.0 * GB + 0.05 * (weights or 0))


def fit_class(total_bytes: int | None, ram_gb: float | None) -> str | None:
    if total_bytes is None or not ram_gb:
        return None
    ratio = total_bytes / usable_gpu_bytes(ram_gb)
    if ratio < 0.70:
        return "Comfortable"
    if ratio < 0.85:
        return "Likely"
    if ratio < 1.0:
        return "Borderline"
    return "Unlikely"


@dataclass
class MemoryEstimate:
    weights_gb: float | None
    weights_source: str
    kv_gb: float | None
    kv_rough: bool
    overhead_gb: float | None
    total_gb: float | None
    context: int
    fit: str | None
    usable_gb: float | None
    exceeds_model_context: bool
    ratio: float | None = None  # total / usable GPU memory

    def to_dict(self) -> dict:
        return asdict(self)


def estimate(
    *,
    params: float | None,
    bits: float | None,
    mode: str | None,
    file_bytes: int | None,
    arch: Arch,
    context: int,
    ram_gb: float | None,
) -> MemoryEstimate:
    w, wsrc = weights_bytes(file_bytes, params, bits, mode)
    kv, upper = kv_cache_bytes(arch, context)
    if kv is None and params:
        # No architecture: rule of thumb between GQA (~0.02) and full multi-head (~0.07)
        # models, in MB of fp16 KV per token per billion params.
        kv = int(params / 1e9 * 0.05 * 1024**2 * context)
        upper = True
    oh = overhead_bytes(w) if w is not None else None
    total = (w + (kv or 0) + oh) if w is not None else None
    r = lambda b: None if b is None else round(b / GB, 2)
    return MemoryEstimate(
        weights_gb=r(w),
        weights_source=wsrc,
        kv_gb=r(kv),
        kv_rough=upper,
        overhead_gb=r(oh),
        total_gb=r(total),
        context=context,
        fit=fit_class(total, ram_gb),
        usable_gb=r(usable_gpu_bytes(ram_gb)) if ram_gb else None,
        exceeds_model_context=bool(arch.max_context and context > arch.max_context),
        ratio=round(total / usable_gpu_bytes(ram_gb), 3) if total is not None and ram_gb else None,
    )