CSIGv3_train_script / src /latency_test.py
XenderYang's picture
CSIGv3 AdcSR train scripts + A100 runbook
4811c23 verified
Raw
History Blame Contribute Delete
1.58 kB
#!/usr/bin/env python
"""512x512 fp16 时延测速: 加载 torch.jit 模型, warmup + N 次前向。
用法: python src/latency_test.py --model model_dir/your_model.pt [--osediff_latency 0.168]
"""
import argparse, time, statistics
import torch
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--model", required=True)
ap.add_argument("--n", type=int, default=100)
ap.add_argument("--warmup", type=int, default=10)
ap.add_argument("--osediff_latency", type=float, default=0.0, help="秒; 提供则打印加速比")
args = ap.parse_args()
device = "cuda" if torch.cuda.is_available() else "cpu"
m = torch.jit.load(args.model, map_location=device)
m.eval()
x = torch.randn(1, 3, 512, 512, device=device).half()
with torch.no_grad():
for _ in range(args.warmup):
m(x)
torch.cuda.synchronize() if device == "cuda" else None
times = []
for _ in range(args.n):
if device == "cuda":
torch.cuda.synchronize()
t0 = time.perf_counter()
with torch.no_grad():
m(x)
if device == "cuda":
torch.cuda.synchronize()
times.append(time.perf_counter() - t0)
mean = statistics.mean(times)
med = statistics.median(times)
print(f"mean {mean*1000:.3f} ms | median {med*1000:.3f} ms | n={args.n}")
if args.osediff_latency > 0:
print(f"speedup vs OSEDiff({args.osediff_latency*1000:.1f}ms): {args.osediff_latency/mean:.2f}x")
if __name__ == "__main__":
main()