File size: 2,359 Bytes
12496fc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Transparent planning arithmetic. FLOPs/memory estimates are not benchmarks."""
from dataclasses import dataclass, asdict


@dataclass
class Estimate:
    total_parameters: float
    active_parameters: float
    tokens: float
    gpus: int
    peak_tflops: float = 989
    mfu: float = 0.35
    gpu_hour_usd: float = 3.0

    def calculate(self):
        if not 0 < self.active_parameters <= self.total_parameters or min(self.tokens, self.gpus, self.peak_tflops, self.gpu_hour_usd) <= 0 or not 0 < self.mfu <= 1:
            raise ValueError("Invalid estimation assumptions")
        flops = 6*self.active_parameters*self.tokens
        hours = flops/(self.gpus*self.peak_tflops*1e12*self.mfu*3600)
        return {**asdict(self), "training_flops": flops, "hours": hours, "gpu_hours": hours*self.gpus,
                "compute_usd": hours*self.gpus*self.gpu_hour_usd,
                "weight_GB": {"bf16": self.total_parameters*2/1e9, "fp8": self.total_parameters/1e9,
                              "int8": self.total_parameters/1e9, "int4_ideal": self.total_parameters/2e9},
                "adam_training_state_GB": self.total_parameters*16/1e9,
                "checkpoint_weights_master_moments_GB": self.total_parameters*14/1e9,
                "token_storage_uint32_GB": self.tokens*4/1e9,
                "limitations": "6NT omits attention quadratic term, routing, communication and recomputation; quantization excludes scales/metadata; costs illustrative"}


def kv_cache_bytes(layers, kv_heads, head_dim, context, batch=1, bytes_per_element=2):
    if min(layers, kv_heads, head_dim, context, batch, bytes_per_element) <= 0:
        raise ValueError("KV dimensions must be positive")
    return 2*layers*kv_heads*head_dim*context*batch*bytes_per_element


def topology(nodes, gpus_per_node, tp, pp, cp, dp, ep=1):
    if min(nodes, gpus_per_node, tp, pp, cp, dp, ep) < 1:
        raise ValueError("Parallel sizes must be positive")
    world = nodes*gpus_per_node
    if tp*pp*cp*dp != world or dp % ep:
        raise ValueError("Require world=TP*PP*CP*DP and EP divides DP (expert subgroup convention)")
    return {"world": world, "tp": tp, "pp": pp, "cp": cp, "dp": dp, "ep": ep, "expert_data_parallel": dp//ep,
            "note": "EP is a subgroup of DP here, not another world-size multiplier; framework-specific mapping must be validated"}