File size: 4,229 Bytes
9f85661
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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:]))