File size: 2,084 Bytes
a64255e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Simulate ZeroGPU locally: load weights on CPU in the parent (CUDA must stay
uninitialized there), then run whole requests in forked children, which is
what ``@spaces.GPU`` does on ZeroGPU.

CUDA_VISIBLE_DEVICES=0 python tools/fork_test.py bagel [n_requests]
"""
import multiprocessing as mp
import sys
import time
from pathlib import Path

sys.path.insert(0, str(Path(__file__).parent.parent))
sys.path.insert(0, str(Path(__file__).parent))

# ZeroGPU's own main-process emulation (fake CUDA queries, CUDA init forbidden); children undo it,
# as spaces' worker does before attaching the GPU.
from spaces.zero.torch import patching  # noqa: E402
import torch  # noqa: E402

patching.patch()
from demo_models import MODELS, run  # noqa: E402
from general_qa import general_qa  # noqa: E402

key = sys.argv[1]
n = int(sys.argv[2]) if len(sys.argv) > 2 else 2
spec = MODELS[key]
t = time.time()
spec.load()
print(f"parent: loaded on CPU in {time.time() - t:.0f}s; CUDA initialized in parent: {torch.cuda.is_initialized()}",
      flush=True)
assert not torch.cuda.is_initialized(), "the main process must not initialize CUDA (ZeroGPU forks it)"
queries = general_qa(per_subset=1, offset=420)[:n]


def child(i, q):
    patching.unpatch()
    t0 = time.time()
    for ev in run(spec, q["image"], q["question"], 60, False):
        if ev[0] == "greedy":
            print(f"  [{i}] greedy={ev[1].answer!r} p0={ev[1].confidence:.2f} k={ev[2]} at {time.time() - t0:.1f}s",
                  flush=True)
        elif ev[0] == "demo":
            print(f"  [{i}] demo [{ev[1].skill}] {ev[1].question!r} -> {ev[1].answer!r} at {time.time() - t0:.1f}s",
                  flush=True)
        elif ev[0] == "final":
            print(f"  [{i}] final={ev[1]!r} {ev[2]} refs={q['answers'][:3]} total {time.time() - t0:.1f}s", flush=True)


ctx = mp.get_context("fork")
for i, q in enumerate(queries):
    p = ctx.Process(target=child, args=(i, q))
    t0 = time.time()
    p.start()
    p.join()
    print(f"request {i} ({q['subset']}): exit={p.exitcode} in {time.time() - t0:.1f}s", flush=True)