import argparse
import os
import subprocess
import sys
import xml.etree.ElementTree as ET
from collections import defaultdict
from statistics import median
XCTRACE = "/Applications/Xcode.app/Contents/Developer/usr/bin/xctrace"
def export_table(trace_path: str, schema: str = "metal-gpu-execution-points") -> str:
if not os.path.isdir(trace_path):
raise FileNotFoundError(f"trace bundle not found: {trace_path}")
xpath = f'/trace-toc/run/data/table[@schema="{schema}"]'
proc = subprocess.run(
[XCTRACE, "export", "--input", trace_path, "--xpath", xpath],
check=True,
capture_output=True,
text=True,
)
return proc.stdout
def resolve_value(elem: ET.Element, dict_table: dict) -> str:
rid = elem.get("id")
ref = elem.get("ref")
if rid is not None:
val = elem.get("fmt") or (elem.text or "")
dict_table[rid] = val
return val
if ref is not None:
return dict_table.get(ref, "")
return elem.get("fmt") or (elem.text or "")
def parse_dispatches(xml_text: str) -> list:
root = ET.fromstring(xml_text)
dispatches = []
dict_table: dict = {}
for row in root.iter("row"):
children = list(row)
if len(children) < 5:
continue
t_str = resolve_value(children[0], dict_table)
t_text = (children[0].text or "").strip()
try:
t_ns = int(t_text) if t_text else int(t_str)
except ValueError:
try:
t_ns = int(t_str)
except ValueError:
continue
chan_str = resolve_value(children[1], dict_table)
fn_str = resolve_value(children[2], dict_table)
try:
fn = int(fn_str)
except ValueError:
continue
slot_str = resolve_value(children[3], dict_table)
sub_str = resolve_value(children[4], dict_table)
dispatches.append({
"t_ns": t_ns,
"channel": chan_str,
"fn": fn,
"slot": slot_str,
"sub_id": sub_str,
})
return dispatches
def pair_dispatches(rows: list) -> list:
starts = {}
paired = []
unpaired_ends = 0
for r in rows:
key = r["sub_id"]
if r["fn"] == 1:
starts[key] = r
elif r["fn"] == 2:
s = starts.pop(key, None)
if s is None:
unpaired_ends += 1
continue
paired.append({
"sub_id": key,
"channel": s["channel"],
"start_ns": s["t_ns"],
"end_ns": r["t_ns"],
"duration_ns": r["t_ns"] - s["t_ns"],
})
return paired, unpaired_ends, len(starts)
def percentile(values: list, p: float) -> float:
if not values:
return 0.0
s = sorted(values)
k = (len(s) - 1) * (p / 100.0)
lo = int(k)
hi = min(lo + 1, len(s) - 1)
frac = k - lo
return s[lo] * (1 - frac) + s[hi] * frac
def summarize(label: str, paired: list) -> dict:
durs = [p["duration_ns"] for p in paired]
if not durs:
return {"label": label, "count": 0}
return {
"label": label,
"count": len(durs),
"sum_ns": sum(durs),
"min_ns": min(durs),
"p50_ns": int(median(durs)),
"p90_ns": int(percentile(durs, 90)),
"p95_ns": int(percentile(durs, 95)),
"p99_ns": int(percentile(durs, 99)),
"max_ns": max(durs),
"mean_ns": int(sum(durs) / len(durs)),
}
def fmt_us(ns: int) -> str:
return f"{ns / 1000.0:.1f}"
def print_summary(s: dict):
if s["count"] == 0:
print(f" {s['label']}: NO DISPATCHES PAIRED")
return
print(f" {s['label']}:")
print(f" count : {s['count']:>8d}")
print(f" sum : {fmt_us(s['sum_ns']):>8s} µs total ({s['sum_ns']/1e9:.3f} s)")
print(f" mean : {fmt_us(s['mean_ns']):>8s} µs/dispatch")
print(f" min : {fmt_us(s['min_ns']):>8s} µs")
print(f" p50 : {fmt_us(s['p50_ns']):>8s} µs")
print(f" p90 : {fmt_us(s['p90_ns']):>8s} µs")
print(f" p95 : {fmt_us(s['p95_ns']):>8s} µs")
print(f" p99 : {fmt_us(s['p99_ns']):>8s} µs")
print(f" max : {fmt_us(s['max_ns']):>8s} µs")
def print_comparison(a: dict, b: dict):
print()
print("Comparison:")
print(f" {'metric':<10s} {a['label']:>20s} {b['label']:>20s} ratio")
print(f" {'-'*10} {'-'*20} {'-'*20} {'-'*5}")
if a["count"] == 0 or b["count"] == 0:
print(" (one side empty — comparison skipped)")
return
for key, label in [
("count", "count"),
("sum_ns", "sum"),
("mean_ns", "mean µs"),
("p50_ns", "p50 µs"),
("p90_ns", "p90 µs"),
("p99_ns", "p99 µs"),
("max_ns", "max µs"),
]:
av, bv = a[key], b[key]
ratio = (av / bv) if bv else float("inf")
if key == "count":
print(f" {label:<10s} {av:>20d} {bv:>20d} {ratio:>5.2f}×")
else:
print(f" {label:<10s} {fmt_us(av):>20s} {fmt_us(bv):>20s} {ratio:>5.2f}×")
def main():
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument("--trace", action="append", required=True,
help="Path to .trace bundle (repeatable; first 2 used for comparison)")
ap.add_argument("--label", action="append", default=[],
help="Label for each --trace (defaults to bundle basename)")
ap.add_argument("--schema", default="metal-gpu-execution-points",
help="xctrace schema to query (default: metal-gpu-execution-points)")
args = ap.parse_args()
summaries = []
for i, t in enumerate(args.trace):
label = args.label[i] if i < len(args.label) else os.path.basename(t).replace(".trace", "")
print(f"=== {label} ({t}) ===")
xml_text = export_table(t, args.schema)
rows = parse_dispatches(xml_text)
paired, unpaired_ends, leftover_starts = pair_dispatches(rows)
print(f" rows parsed : {len(rows)}")
print(f" dispatches paired: {len(paired)}")
print(f" unpaired ends : {unpaired_ends}")
print(f" leftover starts : {leftover_starts}")
s = summarize(label, paired)
print_summary(s)
summaries.append(s)
print()
if len(summaries) >= 2:
print_comparison(summaries[0], summaries[1])
if __name__ == "__main__":
main()