Add thread_flamegraph.py

This is a perfetto trace analysis script that should be invoked upon a
`.pftrace` script that includes callstack sampling, such as one from
`./perfetto.sh scheduling`.  It finds the top ten threads by runtime,
then builds a perf flamegraph and then analyzes scheduler latency for
each thread to see what CPU thread contention looks like.
This commit is contained in:
Elliot Saba committed 2025-10-20 12:26:14 -07:00
1 parent 63bf56e24c
commit 969e8515a6
4 files changed
+333

No files matched your search

+3
View File
@@ -0,0 +1,3 @@
[submodule "perfetto/FlameGraph"]
path = perfetto/FlameGraph
url = https://github.com/brendangregg/FlameGraph
+1
View File
@@ -0,0 +1 @@
*.pftrace
Submodule perfetto/FlameGraph added at 41fee1f99f.
+328
View File
@@ -0,0 +1,328 @@
#!/usr/bin/env python3
"""
thread_flamegraphs.py
Given a Perfetto trace file, produce:
- top 10 threads by CPU time
- per-thread folded stacks and flamegraph SVG (via flamegraph.pl)
- per-thread summary of "waiting" time (sum of thread_state.dur where state != 'Running')
Usage:
python thread_flamegraphs.py TRACE_FILE.pftrace \
--outdir out_dir \
[--top N]
Requires:
- perfetto.trace_processor Python bindings (trace_processor)
- flamegraph.pl (optional; script will still emit .folded files)
"""
import argparse, os, subprocess, sys, importlib
from collections import Counter, defaultdict
def ensure_package(package_name):
"""
Ensures that a given package is installed and imported.
If the import fails, it attempts to install the package using pip with --user.
Returns the imported module object.
"""
try:
return importlib.import_module(package_name)
except ImportError:
print(f"Package '{package_name}' not found. Attempting to install...")
subprocess.check_call([sys.executable, "-m", "pip", "install", "--user", "--break-system-packages", package_name])
print(f"Successfully installed '{package_name}'. Importing...")
ensure_package("matplotlib")
ensure_package("numpy")
ensure_package("perfetto")
# Perfetto TraceProcessor Python API
from perfetto.trace_processor import TraceProcessor
# Helpers for plotting and generating histograms
import matplotlib.pyplot as plt
import numpy as np
# ---------- Helper SQL snippets ----------
SQL_TOP_THREADS = """
SELECT
sched.utid AS utid,
thread.tid AS tid,
thread.name AS thread_name,
process.name AS process_name,
SUM(sched.dur) AS cpu_ns
FROM sched
LEFT JOIN thread USING(utid)
LEFT JOIN process ON thread.upid = process.upid
WHERE NOT thread.is_idle
GROUP BY sched.utid
ORDER BY cpu_ns DESC
LIMIT {limit};
"""
SQL_THREAD_STATE_SUM = """
SELECT state, SUM(dur) AS dur_ns
FROM thread_state
WHERE utid = {utid}
GROUP BY state;
"""
# fetch CPU-profile (perf) samples for a thread (either perf_sample or cpu_profile_stack_sample)
SQL_SAMPLES_BY_UTID = """
SELECT id, ts, callsite_id
FROM cpu_profile_stack_sample
WHERE utid = {utid}
UNION ALL
SELECT id, ts, callsite_id
FROM perf_sample
WHERE utid = {utid}
ORDER BY ts;
"""
# read the callsite table and frame table (we load all of them once)
SQL_STACK_CALLSITES = "SELECT id, parent_id, frame_id FROM stack_profile_callsite;"
SQL_STACK_FRAMES = "SELECT id, name FROM stack_profile_frame;"
# ---------- Main logic ----------
def make_folded_from_callsites(callsite_ids, callsite_map, frame_map):
"""
callsite_ids: iterable of callsite_id ints (may include None)
callsite_map: dict id -> (parent_id, frame_id)
frame_map: dict frame_id -> name
Return: Counter mapping folded_stack_string -> count
"""
folded = Counter()
for cs in callsite_ids:
if cs is None:
# Some samples might not have a callsite (NULL). skip/mark.
folded["(no_stack)"] += 1
continue
stack_frames = []
cur = cs
# walk parent chain bottom->top, but stack_profile_callsite parent points to parent node;
# frame_id is the frame for that callsite node
while cur is not None:
entry = callsite_map.get(cur)
if entry is None:
# defensive: missing callsite
break
parent_id, frame_id = entry
# frame name lookup
fname = frame_map.get(frame_id, "<unknown>")
# We build stack bottom -> top, but folded needs topmost last -> we will reverse later
stack_frames.append(fname)
cur = parent_id
# 'stack_frames' currently bottom-most first (bottom = leaf?). The exact order may vary depending
# on how callsite parent is defined; empirically we want "root;...;leaf" for flamegraph,
# so reverse to get root->leaf (so flamegraph shows root at bottom).
stack_frames = list(reversed(stack_frames))
if not stack_frames:
folded["(empty)"] += 1
else:
folded_key = ";".join(stack_frames)
folded[folded_key] += 1
return folded
def run_flamegraph_pl(folded_file, out_svg, flamegraph_pl_path="FlameGraph/flamegraph.pl"):
"""Run flamegraph.pl folded->svg. Returns (ok, output)"""
cmd = f"c++filt < {folded_file} | {flamegraph_pl_path} > {out_svg}"
# flamegraph.pl writes svg to stdout by default; redirect
try:
p = subprocess.run(cmd, shell=True, check=False)
if p.returncode != 0:
return False, p.stderr.decode("utf-8", errors="replace")
return True, None
except FileNotFoundError:
return False, f"flamegraph.pl not found at '{flamegraph_pl_path}'."
except Exception as e:
return False, str(e)
def plot_waiting_histogram(tp, tid, utid, out_svg, bins='auto', logx=True, logy=True):
"""
Generate an SVG histogram of waiting durations for a given thread (utid).
Returns basic stats (count, mean, median, p95).
"""
q = f"SELECT dur FROM thread_state WHERE utid={utid} AND state!='Running';"
durs = [int(r.dur) for r in tp.query(q) if getattr(r, "dur", None)]
if not durs:
print(f" No waiting intervals for utid={utid}")
return None
# Convert ns → ms for readability
durs_ms = np.array(durs, dtype=np.float64) / 1e6
count = len(durs_ms)
mean = float(np.mean(durs_ms))
median = float(np.median(durs_ms))
p95 = float(np.percentile(durs_ms, 95))
# Choose log bins for wide distributions
if logx:
min_v = max(1e-3, durs_ms.min())
max_v = durs_ms.max()
bins = np.logspace(np.log10(min_v), np.log10(max_v), 50)
else:
bins = 'auto'
plt.figure(figsize=(6, 4))
plt.hist(durs_ms, bins=bins, color='gray', alpha=0.8)
plt.xlabel("Wait duration (ms)")
plt.ylabel("Count")
plt.title(f"Thread {tid} waiting time histogram\nmedian={median:.3f} ms, p95={p95:.3f} ms")
if logx:
plt.xscale('log')
if logy:
plt.yscale('log')
plt.grid(True, which='both', ls=':')
plt.tight_layout()
plt.savefig(out_svg, format='svg')
plt.close()
print(f" Wrote waiting-time histogram to {out_svg}")
return {"count": count, "mean_ms": mean, "median_ms": median, "p95_ms": p95}
def main():
ap = argparse.ArgumentParser(description="Perfetto trace -> top-thread flamegraphs & waiting time.")
ap.add_argument("trace", help="Perfetto trace file (e.g. trace.perfetto-trace)")
ap.add_argument("--top", type=int, default=10, help="number of top threads to analyze (default 10)")
args = ap.parse_args()
trace_file = args.trace
outdir = os.path.basename(trace_file)
top_n = args.top
os.makedirs(outdir, exist_ok=True)
print(f"Opening trace: {trace_file}")
tp = TraceProcessor(trace=trace_file)
# 1) get top threads by CPU time (sched.dur)
q = SQL_TOP_THREADS.format(limit=top_n)
print("Querying top threads by SUM(sched.dur)...")
rows = list(tp.query(q))
if not rows:
print("No sched rows found in trace. Exiting.")
return 1
top_threads = []
print("Top threads (utid, name, process, cpu_ms):")
for r in rows:
# r is a row-like object (namedtuple/dict-like depending on the binding)
utid = int(r.utid)
tid = int(r.tid)
tname = r.thread_name if getattr(r, "thread_name", None) else "<unnamed>"
pname = r.process_name if getattr(r, "process_name", None) else "<no-process>"
cpu_ns = int(r.cpu_ns) if getattr(r, "cpu_ns", None) else 0
cpu_ms = cpu_ns / 1e6
print(f" utid={utid} thread='{tname}' process='{pname}' cpu_ms={cpu_ms:.3f}")
top_threads.append((utid, tid, tname, pname, cpu_ns))
# Preload stack callsite/frame tables (to reconstruct stacks)
print("Loading stack_profile_callsite and stack_profile_frame tables (to reconstruct callstacks)...")
callsite_map = {} # id -> (parent_id, frame_id)
try:
for r in tp.query(SQL_STACK_CALLSITES):
callsite_map[int(r.id)] = (None if r.parent_id is None else int(r.parent_id),
None if r.frame_id is None else int(r.frame_id))
except Exception as e:
# If the table doesn't exist, we still continue (maybe samples are not present)
print("Warning: could not load stack_profile_callsite table:", e)
frame_map = {} # frame_id -> name
try:
for r in tp.query(SQL_STACK_FRAMES):
frame_map[int(r.id)] = (r.name or "<no_name>")
except Exception as e:
print("Warning: could not load stack_profile_frame table:", e)
# For each top thread, get samples and build folded stacks
results = []
for idx, (utid, tid, tname, pname, cpu_ns) in enumerate(top_threads, start=1):
print(f"\n[{idx}/{len(top_threads)}] Processing tid={tid} thread='{tname}'")
# collect sample callsite ids
callsite_ids = []
try:
for r in tp.query(SQL_SAMPLES_BY_UTID.format(utid=utid)):
# callsite_id may be None
callsite_ids.append(None if getattr(r, "callsite_id", None) is None else int(r.callsite_id))
except Exception as e:
# Samples tables might differ in traces; try querying cpu_profile_stack_sample only
print(" Warning: union query failed; trying cpu_profile_stack_sample only...", e)
try:
q2 = "SELECT id, ts, callsite_id FROM cpu_profile_stack_sample WHERE utid = {utid} ORDER BY ts;".format(utid=utid)
for r in tp.query(q2):
callsite_ids.append(None if getattr(r, "callsite_id", None) is None else int(r.callsite_id))
except Exception as e2:
print(" No samples available for utid", utid, ":", e2)
callsite_ids = []
print(f" Collected {len(callsite_ids)} samples (callsite entries).")
# Build folded stacks from these callsite ids
folded = make_folded_from_callsites(callsite_ids, callsite_map, frame_map)
folded_file = os.path.join(outdir, f"{idx}_thread_{tid}_{sanitize_filename(tname)}.folded")
svg_file = os.path.join(outdir, f"{idx}_thread_{tid}_{sanitize_filename(tname)}.svg")
# write folded file
with open(folded_file, "w", encoding="utf-8") as f:
for stack, count in folded.most_common():
f.write(f"{stack} {count}\n")
print(f" Wrote folded stacks to {folded_file} ({len(folded)} unique stacks).")
ok, err = run_flamegraph_pl(folded_file, svg_file)
if ok:
print(f" Generated flamegraph SVG: {svg_file}")
else:
print(f" Could not run flamegraph.pl: {err}")
# compute waiting time using thread_state table (sum of durations for non-Running states)
hist_svg = os.path.join(outdir, f"{idx}_thread_{tid}_{sanitize_filename(tname)}_waiting_hist.svg")
hist_stats = plot_waiting_histogram(tp, tid, utid, hist_svg)
# total_waiting_ns = 0
# try:
# # get per-state sums
# qts = SQL_THREAD_STATE_SUM.format(utid=utid)
# for r in tp.query(qts):
# state = r.state
# dur_ns = int(r.dur_ns) if getattr(r, "dur_ns", None) is not None else int(r.dur)
# # treat 'Running' as not-waiting; everything else as waiting (adjust if you prefer)
# if state != "Running":
# total_waiting_ns += dur_ns
# except Exception as e:
# print(" Warning: thread_state query failed:", e)
results.append({
"utid": utid,
"thread_name": tname,
"process_name": pname,
"cpu_ns": cpu_ns,
"samples": len(callsite_ids),
"unique_stacks": len(folded),
"folded_file": folded_file,
"svg_file": svg_file,
"waiting_ms_median": hist_stats['median_ms'],
"waiting_ms_p95": hist_stats['p95_ms'],
})
# Summary output
print("\n=== Summary for top threads ===")
for r in results:
cpu_ms = r["cpu_ns"] / 1e6
print(f"utid={r['utid']} thread='{r['thread_name']}' proc='{r['process_name']}' cpu_ms={cpu_ms:.3f} wait_ms=({r["waiting_ms_median"]:.3f}, @95: {r["waiting_ms_p95"]:.3f}) samples={r['samples']} unique_stacks={r['unique_stacks']}")
print(f" folded: {r['folded_file']}")
if r['svg_file']:
print(f" svg: {r['svg_file']}")
print("\nDone.")
return 0
def sanitize_filename(s):
# simple sanitizer for thread names in filenames
out = "".join(c if (c.isalnum() or c in "._-") else "_" for c in (s or "thread"))
return out[:120]
if __name__ == "__main__":
sys.exit(main())