simit / tools /fork_test.py
monurcan's picture
SIMIT demo: BAGEL, Lance, Qwen3.8-27B-FP8 + FLUX.2-klein on ZeroGPU
a64255e verified
Raw History Blame Contribute Delete
2.08 kB
"""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)