distinct / scripts /offload_probe.py
User1342's picture
Measure the GPU offload instead of estimating it, and keep the development scratch out of the repository
9f85661
Raw History Blame Contribute Delete
4.23 kB
"""Time this machine at several GPU offloads and print what each achieved.
The worker does this once by itself and writes the answer down. This is the
same measurement with the table printed, for when somebody wants to see the
numbers rather than trust them -- or to re-take them after a driver update.
python scripts/offload_probe.py # measure, or show what is cached
python scripts/offload_probe.py --force # measure again regardless
"""
from __future__ import annotations
import argparse
import sys
import time
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from distinct_agent import offload as offload_module # noqa: E402
from distinct_agent import runtime # noqa: E402
from distinct_agent.gguf import read_shape # noqa: E402
from distinct_agent.server_runner import LlamaServerRunner, plan_offload # noqa: E402
from distinct_agent.weights import ( # noqa: E402
DEFAULT_CACHE_GB,
WeightsCache,
cache_bytes,
fetchable_manifests,
shared_cache_directory,
)
MB = 1024 * 1024
def main(argv: list[str]) -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model", default="", help="defaults to the first cached model")
parser.add_argument("--llama-server", default="")
parser.add_argument("--context", type=int, default=4096)
parser.add_argument("--force", action="store_true", help="ignore anything cached")
args = parser.parse_args(argv)
cache = WeightsCache(
shared_cache_directory(),
limit_bytes=cache_bytes(DEFAULT_CACHE_GB),
manifests=tuple(fetchable_manifests()),
)
installed = list(cache.installed())
if args.model:
installed = [item for item in installed if item.manifest.id == args.model]
if not installed:
print("no cached model to measure with", file=sys.stderr)
return 1
model = installed[0]
executable = args.llama_server or runtime.best_installed()
if not executable:
print("no llama-server installed", file=sys.stderr)
return 1
free, total = runtime.accelerator_memory()
shape = read_shape(model.path)
size = model.path.stat().st_size
plan = plan_offload(
model_bytes=size,
shape=shape,
free_vram_bytes=free,
context_tokens=args.context,
)
settings = offload_module.candidates(plan)
print(f"gpu {runtime.accelerator_name() or 'none'}")
print(f"vram {free / MB:.0f} MB free of {total / MB:.0f} MB")
print(f"build {Path(executable).parent.name}")
print(f"model {model.manifest.id} ({size / MB:.0f} MB, {shape.layers} layers)")
print(f"fits {plan}")
print(f"timing {settings}\n")
key = offload_module.machine_key(
model_id=model.manifest.id,
gpu=runtime.accelerator_name(),
vram_bytes=total,
build=Path(executable).parent.name,
context=args.context,
)
if not args.force:
known = offload_module.remembered(key)
if known is not None:
print(f"already measured: {known} layer(s). Pass --force to measure again.")
return 0
runner = LlamaServerRunner(executable, context_length=args.context)
runner._notify = lambda message: print(f" {message}", flush=True)
started = time.monotonic()
with runner._lock:
measurements = runner._time_offloads(model, settings)
chosen = offload_module.choose(measurements)
offload_module.remember(key, chosen, measurements)
print(f"\n{'layers':>7} {'tok/s':>8} note")
print("-" * 34)
for item in measurements:
print(f"{item.layers:>7} {item.tokens_per_second:>8.2f} {item.note}")
print("-" * 34)
baseline = next((m.tokens_per_second for m in measurements if m.layers == 0), 0.0)
best = max((m.tokens_per_second for m in measurements), default=0.0)
if baseline > 0:
print(f"best is {best / baseline:.2f}x the CPU")
print(f"chosen: {chosen} layer(s) on the GPU")
print(f"measured in {time.monotonic() - started:.0f}s, written to {offload_module.cache_path()}")
return 0
if __name__ == "__main__":
raise SystemExit(main(sys.argv[1:]))