InferScale-Sim / py /inferscale /latency.py
ArchitSharma's picture
Deepen InferScale simulation research workflow
ce2d64b
Raw
History Blame Contribute Delete
4.7 kB
from __future__ import annotations
import math
from .models import AcceleratorProfile, ModelProfile
from .profiles import QUANTIZATION_BYTES, QUANTIZATION_COMPUTE_MULTIPLIER
class AnalyticalLatencyModel:
"""Hardware-aware analytical proxy for prefill/decode latency.
This model intentionally does *not* claim benchmark-grade accuracy. It uses
model architecture, accelerator peak specs, conservative efficiency factors,
and a roofline-style max(compute_time, memory_time) approximation. Public
profiles are tagged `analytical-reference` throughout the app.
"""
def __init__(
self,
model: ModelProfile,
accelerator: AcceleratorProfile,
quantization: str = "fp16",
prefill_scale: float = 1.0,
decode_scale: float = 1.0,
):
if quantization not in QUANTIZATION_BYTES:
raise ValueError(f"Unsupported quantization: {quantization}")
self.model = model
self.accelerator = accelerator
self.quantization = quantization
self.weight_bytes_per_param = QUANTIZATION_BYTES[quantization]
self.compute_overhead = QUANTIZATION_COMPUTE_MULTIPLIER[quantization]
self.prefill_scale = max(float(prefill_scale), 1e-6)
self.decode_scale = max(float(decode_scale), 1e-6)
@property
def model_weight_gb(self) -> float:
return self.model.params_b * self.weight_bytes_per_param
def kv_bytes_per_token(self) -> float:
# K + V, all layers, KV heads only. KV state is assumed fp16 in the current model.
return (
2
* self.model.layers
* self.model.kv_heads
* self.model.head_dim
* 2.0
)
def _compute_efficiency(self, batch_size: int, tokens: int) -> float:
scale = 1.0 + 0.12 * math.log2(max(batch_size, 1)) + 0.035 * math.log2(max(tokens, 1))
return min(0.88, self.accelerator.compute_efficiency * scale)
def _bandwidth_efficiency(self, batch_size: int) -> float:
scale = 1.0 + 0.06 * math.log2(max(batch_size, 1))
return min(0.91, self.accelerator.bandwidth_efficiency * scale)
def prefill_seconds(self, token_counts: list[int]) -> float:
if not token_counts:
return 0.0
batch = len(token_counts)
total_tokens = sum(token_counts)
max_seq = max(token_counts)
dense_flops = 2.0 * self.model.params_b * 1e9 * total_tokens
# Approximate quadratic attention component. It is small for short
# contexts but becomes visible at long prompts.
attention_flops = (
4.0
* self.model.layers
* self.model.hidden_size
* sum(t * t for t in token_counts)
)
flops = (dense_flops + attention_flops) * self.compute_overhead
compute = flops / (
self.accelerator.peak_tflops_fp16 * 1e12 * self._compute_efficiency(batch, max_seq)
)
weight_bytes = self.model.params_b * 1e9 * self.weight_bytes_per_param
activation_bytes = total_tokens * self.model.hidden_size * self.model.layers * 2.0 * 0.18
memory = (weight_bytes + activation_bytes) / (
self.accelerator.bandwidth_gbps * 1e9 * self._bandwidth_efficiency(batch)
)
# Kernel launch / scheduling proxy prevents implausibly tiny times.
launch = 0.0018 + 0.00008 * batch
return (max(compute, memory * 0.28) + launch) * self.prefill_scale
def decode_step_seconds(self, context_lengths: list[int]) -> float:
if not context_lengths:
return 0.0
batch = len(context_lengths)
avg_context = sum(context_lengths) / batch
dense_flops = 2.0 * self.model.params_b * 1e9 * batch
attention_flops = (
4.0
* self.model.layers
* self.model.hidden_size
* sum(context_lengths)
)
flops = (dense_flops + attention_flops) * self.compute_overhead
compute = flops / (
self.accelerator.peak_tflops_fp16 * 1e12 * self._compute_efficiency(batch, 1)
)
# Decode is commonly memory-bound. We model one shared weight stream plus
# KV reads that scale with batch and context length.
weight_bytes = self.model.params_b * 1e9 * self.weight_bytes_per_param
kv_read_bytes = self.kv_bytes_per_token() * sum(context_lengths)
memory = (weight_bytes + kv_read_bytes) / (
self.accelerator.bandwidth_gbps * 1e9 * self._bandwidth_efficiency(batch)
)
launch = 0.0012 + 0.000035 * batch + 0.00000003 * avg_context
return (max(compute, memory) + launch) * self.decode_scale