Skip to content

latency

shrinkai.profiler.latency

Functions:

Name Description
measure_latency

Measures average inference latency and throughput (FPS / samples per second).

synchronize_device

Synchronizes device streams to ensure precise time measurement.

Functions:

measure_latency

measure_latency(
    model: Module,
    sample_input: Tensor,
    device: device,
    num_runs: int = 100,
    warmup_runs: int = 15,
) -> dict[str, float]

Measures average inference latency and throughput (FPS / samples per second).

Parameters:

Name Type Description Default
model Module

Model to profile.

required
sample_input Tensor

Input batch tensor matching inference shape.

required
device device

Computing device.

required
num_runs int

Number of timed inference iterations.

100
warmup_runs int

Initial iterations discarded to warm up hardware caches.

15

Returns:

Type Description
dict[str, float]

dict[str, float]: Latency per sample (ms), per batch (ms), and throughput (FPS).

Source code in src/shrinkai/profiler/latency.py
def measure_latency(
    model: nn.Module,
    sample_input: torch.Tensor,
    device: torch.device,
    num_runs: int = 100,
    warmup_runs: int = 15,
) -> dict[str, float]:
    """Measures average inference latency and throughput (FPS / samples per second).

    Args:
        model: Model to profile.
        sample_input: Input batch tensor matching inference shape.
        device: Computing device.
        num_runs: Number of timed inference iterations.
        warmup_runs: Initial iterations discarded to warm up hardware caches.

    Returns:
        dict[str, float]: Latency per sample (ms), per batch (ms), and throughput (FPS).
    """
    was_training = model.training
    first_param = next(model.parameters(), None)
    original_device = first_param.device if first_param is not None else None

    model.eval()
    model.to(device)
    sample_input = sample_input.to(device)
    batch_size = sample_input.size(0)

    with torch.no_grad():
        for _ in range(warmup_runs):
            _ = model(sample_input)
        synchronize_device(device)

    timings: list[float] = []
    with torch.no_grad():
        for _ in range(num_runs):
            synchronize_device(device)
            start_time = time.perf_counter()
            _ = model(sample_input)
            synchronize_device(device)
            timings.append(time.perf_counter() - start_time)

    avg_batch_latency_s = sum(timings) / len(timings)
    avg_batch_latency_ms = avg_batch_latency_s * 1000.0
    avg_sample_latency_ms = avg_batch_latency_ms / batch_size
    fps = (batch_size * num_runs) / sum(timings)

    model.train(mode=was_training)
    if original_device is not None:
        model.to(original_device)

    return {
        "batch_latency_ms": avg_batch_latency_ms,
        "sample_latency_ms": avg_sample_latency_ms,
        "fps": fps,
    }

synchronize_device

synchronize_device(device: device) -> None

Synchronizes device streams to ensure precise time measurement.

Parameters:

Name Type Description Default
device device

Active PyTorch computing device.

required
Source code in src/shrinkai/profiler/latency.py
def synchronize_device(device: torch.device) -> None:
    """Synchronizes device streams to ensure precise time measurement.

    Args:
        device: Active PyTorch computing device.
    """
    if device.type == "cuda":
        torch.cuda.synchronize(device)
    elif device.type == "mps" and hasattr(torch.mps, "synchronize"):
        torch.mps.synchronize()