yolo26-face / scripts /isolate_bottleneck.py
a-ml's picture
Add YOLO26 face-parsing models: PyTorch checkpoints, Core ML exports, training/export scripts, results and demo
e2f3b24 verified
Raw
History Blame Contribute Delete
6 kB
"""
Decisive isolation benchmark: WHERE do the 8.5 s/iteration actually go?
Reported inference is 1.5 ms/image, so a batch-32 fwd+bwd should cost ~150 ms,
yet a training iteration measures ~8500 ms (55x gap). Rather than guess, time
each stage separately:
A. compute -- fwd+bwd on synthetic tensors, no dataloader, no loss
B. compute+loss -- adds the real semantic criterion (scatter ops on MPS)
C. determinism -- same as B with use_deterministic_algorithms(False)
D. dataload -- real dataset pipeline only, no model (workers 0 vs 8)
E. scaling -- fwd+bwd cost for n / l / x variants
Note: another GPU job may be running; treat absolute numbers as contended and
compare stages relative to each other within this run.
"""
import argparse, time
import numpy as np
import torch
def bench(fn, n=8, warmup=3):
for _ in range(warmup):
fn()
torch.mps.synchronize() if torch.backends.mps.is_available() else None
t0 = time.time()
for _ in range(n):
fn()
if torch.backends.mps.is_available():
torch.mps.synchronize()
return (time.time() - t0) / n * 1000
def make_model(name, device):
from ultralytics import YOLO
m = YOLO(name).model.to(device).train()
for p in m.parameters(): # ultralytics freezes params on checkpoint load
p.requires_grad_(True)
return m
def stage_compute(device, batch, imgsz, name="yolo26n-sem.pt"):
m = make_model(name, device)
opt = torch.optim.AdamW(m.parameters(), lr=1e-4)
x = torch.rand(batch, 3, imgsz, imgsz, device=device)
def collect(z, out):
if torch.is_tensor(z):
if z.requires_grad and z.is_floating_point():
out.append(z.float().pow(2).mean())
elif isinstance(z, (list, tuple)):
for e in z:
collect(e, out)
elif isinstance(z, dict):
for e in z.values():
collect(e, out)
def step():
opt.zero_grad(set_to_none=True)
terms = []
collect(m(x), terms)
if not terms:
raise RuntimeError("no differentiable output from model forward")
torch.stack(terms).sum().backward() # trivial surrogate loss
opt.step()
return bench(step)
def stage_compute_loss(device, batch, imgsz, deterministic, name="yolo26n-sem.pt"):
torch.use_deterministic_algorithms(deterministic, warn_only=True)
from ultralytics import YOLO
y = YOLO(name)
m = y.model.to(device).train()
for p in m.parameters():
p.requires_grad_(True)
opt = torch.optim.AdamW(m.parameters(), lr=1e-4)
x = torch.rand(batch, 3, imgsz, imgsz, device=device)
masks = torch.randint(0, 19, (batch, imgsz, imgsz), device=device, dtype=torch.long)
b = {"img": x, "semantic_mask": masks}
def step():
opt.zero_grad(set_to_none=True)
loss, _ = m.loss(b)
(loss.sum() if loss.dim() else loss).backward()
opt.step()
try:
ms = bench(step, n=6, warmup=2)
except Exception as e:
ms = float("nan")
print(" (loss path failed:", repr(e)[:160], ")")
torch.use_deterministic_algorithms(False, warn_only=True)
return ms
def stage_dataload(workers, batch, imgsz, n_batches=8):
from ultralytics.data.dataset import SemanticDataset
from ultralytics.cfg import get_cfg
from ultralytics.utils import DEFAULT_CFG
from torch.utils.data import DataLoader
import yaml
dpath = "/Users/ari/FaceSegmentation/dataset_celebamaskhq_semantic"
with open(f"{dpath}/data.yaml") as f:
data = yaml.safe_load(f)
data["nc"] = len(data["names"])
data["channels"] = 3
data["path"] = dpath
args = get_cfg(DEFAULT_CFG)
args.imgsz = imgsz
args.mosaic = 0.5
args.fliplr = 0.0
args.degrees = 10.0
args.scale = 0.5
ds = SemanticDataset(img_path=f"{dpath}/images/train", imgsz=imgsz, batch_size=batch,
augment=True, hyp=args, data=data, task="semantic")
dl = DataLoader(ds, batch_size=batch, shuffle=True, num_workers=workers,
collate_fn=getattr(ds, "collate_fn", None),
persistent_workers=workers > 0, prefetch_factor=4 if workers > 0 else None)
it = iter(dl)
next(it) # warm
t0 = time.time()
got = 0
for _ in range(n_batches):
try:
next(it)
got += 1
except StopIteration:
break
return (time.time() - t0) / max(got, 1) * 1000
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--batch", type=int, default=32)
ap.add_argument("--imgsz", type=int, default=512)
ap.add_argument("--device", default="mps")
ap.add_argument("--stages", default="A,B,C,D,E")
args = ap.parse_args()
S = set(args.stages.split(","))
dev = args.device
B, R = args.batch, args.imgsz
print(f"device={dev} batch={B} imgsz={R}\n")
if "A" in S:
ms = stage_compute(dev, B, R)
print(f"A. compute only (fwd+bwd, synthetic, no loss) : {ms:8.0f} ms/iter")
if "B" in S:
ms = stage_compute_loss(dev, B, R, deterministic=True)
print(f"B. compute + real loss, deterministic=True : {ms:8.0f} ms/iter")
if "C" in S:
ms = stage_compute_loss(dev, B, R, deterministic=False)
print(f"C. compute + real loss, deterministic=False : {ms:8.0f} ms/iter")
if "D" in S:
for w in (0, 8):
ms = stage_dataload(w, B, R)
print(f"D. dataload only, workers={w} : {ms:8.0f} ms/batch")
if "E" in S:
for nm in ("yolo26n-sem.pt", "yolo26s-sem.pt", "yolo26m-sem.pt", "yolo26l-sem.pt", "yolo26x-sem.pt"):
try:
ms = stage_compute(dev, B, R, nm)
print(f"E. compute {nm:16s} : {ms:8.0f} ms/iter")
except Exception as e:
print(f"E. compute {nm:16s} : FAILED {repr(e)[:90]}")
if __name__ == "__main__":
main()