[object Object]

← back to Exo

Ciaran/profiling (#1345)

19965c7ba54dcab1f6446e01e87f9a8610dcce7c · 2026-02-02 17:19:45 +0000 · ciaranbor

## Motivation

Know what the hell is going on

## Changes

- Tracing library (src/exo/shared/tracing.py): trace() context manager,
Chrome Trace Format export, statistics computation
- Runner instrumentation
(src/exo/worker/engines/image/pipeline/runner.py): Wrapped sync/async
steps, compute blocks, and send/recv operations
- Trace collection: Workers send traces to master after task completion;
merged into ~/.exo/traces/trace_{task_id}.json
  - API endpoints: List, fetch, stats, and raw download at /v1/traces/*
  - Dashboard: Trace list and detail pages with Perfetto integration

## Why It Works

<img width="1236" height="767" alt="Screenshot 2026-01-30 at 19 00 09"
src="https://github.com/user-attachments/assets/73e6e46d-ba10-4e83-ba99-ff1c3f62ab05"
/>
<img width="1659" height="89" alt="Screenshot 2026-01-30 at 19 00 58"
src="https://github.com/user-attachments/assets/c0fd0e65-e4fc-4fd5-920d-b43b2887d591"
/>

Files touched

Diff

commit 19965c7ba54dcab1f6446e01e87f9a8610dcce7c
Author: ciaranbor <81697641+ciaranbor@users.noreply.github.com>
Date:   Mon Feb 2 17:19:45 2026 +0000

    Ciaran/profiling (#1345)
    
    ## Motivation
    
    Know what the hell is going on
    
    ## Changes
    
    - Tracing library (src/exo/shared/tracing.py): trace() context manager,
    Chrome Trace Format export, statistics computation
    - Runner instrumentation
    (src/exo/worker/engines/image/pipeline/runner.py): Wrapped sync/async
    steps, compute blocks, and send/recv operations
    - Trace collection: Workers send traces to master after task completion;
    merged into ~/.exo/traces/trace_{task_id}.json
      - API endpoints: List, fetch, stats, and raw download at /v1/traces/*
      - Dashboard: Trace list and detail pages with Perfetto integration
    
    ## Why It Works
    
    <img width="1236" height="767" alt="Screenshot 2026-01-30 at 19 00 09"
    src="https://github.com/user-attachments/assets/73e6e46d-ba10-4e83-ba99-ff1c3f62ab05"
    />
    <img width="1659" height="89" alt="Screenshot 2026-01-30 at 19 00 58"
    src="https://github.com/user-attachments/assets/c0fd0e65-e4fc-4fd5-920d-b43b2887d591"
    />
---
 dashboard/src/lib/stores/app.svelte.ts            |  82 +++++
 dashboard/src/routes/traces/+page.svelte          | 190 +++++++++++
 dashboard/src/routes/traces/[taskId]/+page.svelte | 367 ++++++++++++++++++++++
 src/exo/master/api.py                             | 134 ++++++++
 src/exo/master/main.py                            |  64 +++-
 src/exo/shared/apply.py                           |   9 +-
 src/exo/shared/constants.py                       |   3 +
 src/exo/shared/tracing.py                         | 238 ++++++++++++++
 src/exo/shared/types/api.py                       |  42 +++
 src/exo/shared/types/events.py                    |  27 +-
 src/exo/worker/engines/image/pipeline/runner.py   | 334 ++++++++++++--------
 src/exo/worker/runner/runner.py                   |  43 ++-
 12 files changed, 1403 insertions(+), 130 deletions(-)

diff --git a/dashboard/src/lib/stores/app.svelte.ts b/dashboard/src/lib/stores/app.svelte.ts
index 9fa371fb..51de6c66 100644
--- a/dashboard/src/lib/stores/app.svelte.ts
+++ b/dashboard/src/lib/stores/app.svelte.ts
@@ -178,6 +178,36 @@ interface ImageApiResponse {
   data: Array<{ b64_json?: string; url?: string }>;
 }
 
+// Trace API response types
+export interface TraceCategoryStats {
+  totalUs: number;
+  count: number;
+  minUs: number;
+  maxUs: number;
+  avgUs: number;
+}
+
+export interface TraceRankStats {
+  byCategory: Record<string, TraceCategoryStats>;
+}
+
+export interface TraceStatsResponse {
+  taskId: string;
+  totalWallTimeUs: number;
+  byCategory: Record<string, TraceCategoryStats>;
+  byRank: Record<number, TraceRankStats>;
+}
+
+export interface TraceListItem {
+  taskId: string;
+  createdAt: string;
+  fileSize: number;
+}
+
+export interface TraceListResponse {
+  traces: TraceListItem[];
+}
+
 interface RawStateResponse {
   topology?: RawTopology;
   instances?: Record<
@@ -2555,6 +2585,49 @@ class AppStore {
       throw error;
     }
   }
+
+  /**
+   * List all available traces
+   */
+  async listTraces(): Promise<TraceListResponse> {
+    const response = await fetch("/v1/traces");
+    if (!response.ok) {
+      throw new Error(`Failed to list traces: ${response.status}`);
+    }
+    return (await response.json()) as TraceListResponse;
+  }
+
+  /**
+   * Check if a trace exists for a given task ID
+   */
+  async checkTraceExists(taskId: string): Promise<boolean> {
+    try {
+      const response = await fetch(`/v1/traces/${encodeURIComponent(taskId)}`);
+      return response.ok;
+    } catch {
+      return false;
+    }
+  }
+
+  /**
+   * Get computed statistics for a task's trace
+   */
+  async fetchTraceStats(taskId: string): Promise<TraceStatsResponse> {
+    const response = await fetch(
+      `/v1/traces/${encodeURIComponent(taskId)}/stats`,
+    );
+    if (!response.ok) {
+      throw new Error(`Failed to fetch trace stats: ${response.status}`);
+    }
+    return (await response.json()) as TraceStatsResponse;
+  }
+
+  /**
+   * Get the URL for the raw trace file (for Perfetto)
+   */
+  getTraceRawUrl(taskId: string): string {
+    return `/v1/traces/${encodeURIComponent(taskId)}/raw`;
+  }
 }
 
 export const appStore = new AppStore();
@@ -2666,3 +2739,12 @@ export const startDownload = (nodeId: string, shardMetadata: object) =>
   appStore.startDownload(nodeId, shardMetadata);
 export const deleteDownload = (nodeId: string, modelId: string) =>
   appStore.deleteDownload(nodeId, modelId);
+
+// Trace actions
+export const listTraces = () => appStore.listTraces();
+export const checkTraceExists = (taskId: string) =>
+  appStore.checkTraceExists(taskId);
+export const fetchTraceStats = (taskId: string) =>
+  appStore.fetchTraceStats(taskId);
+export const getTraceRawUrl = (taskId: string) =>
+  appStore.getTraceRawUrl(taskId);
diff --git a/dashboard/src/routes/traces/+page.svelte b/dashboard/src/routes/traces/+page.svelte
new file mode 100644
index 00000000..dae70ff5
--- /dev/null
+++ b/dashboard/src/routes/traces/+page.svelte
@@ -0,0 +1,190 @@
+<script lang="ts">
+  import { onMount } from "svelte";
+  import {
+    listTraces,
+    getTraceRawUrl,
+    type TraceListItem,
+  } from "$lib/stores/app.svelte";
+  import HeaderNav from "$lib/components/HeaderNav.svelte";
+
+  let traces = $state<TraceListItem[]>([]);
+  let loading = $state(true);
+  let error = $state<string | null>(null);
+
+  function formatBytes(bytes: number): string {
+    if (!bytes || bytes <= 0) return "0B";
+    const units = ["B", "KB", "MB", "GB"];
+    const i = Math.min(
+      Math.floor(Math.log(bytes) / Math.log(1024)),
+      units.length - 1,
+    );
+    const val = bytes / Math.pow(1024, i);
+    return `${val.toFixed(val >= 10 ? 0 : 1)}${units[i]}`;
+  }
+
+  function formatDate(isoString: string): string {
+    const date = new Date(isoString);
+    return date.toLocaleString();
+  }
+
+  async function downloadTrace(taskId: string) {
+    const response = await fetch(getTraceRawUrl(taskId));
+    const blob = await response.blob();
+    const url = URL.createObjectURL(blob);
+    const a = document.createElement("a");
+    a.href = url;
+    a.download = `trace_${taskId}.json`;
+    a.click();
+    URL.revokeObjectURL(url);
+  }
+
+  async function openInPerfetto(taskId: string) {
+    // Fetch trace data from our local API
+    const response = await fetch(getTraceRawUrl(taskId));
+    const traceData = await response.arrayBuffer();
+
+    // Open Perfetto UI
+    const perfettoWindow = window.open("https://ui.perfetto.dev");
+    if (!perfettoWindow) {
+      alert("Failed to open Perfetto. Please allow popups.");
+      return;
+    }
+
+    // Wait for Perfetto to be ready, then send trace via postMessage
+    const onMessage = (e: MessageEvent) => {
+      if (e.data === "PONG") {
+        window.removeEventListener("message", onMessage);
+        perfettoWindow.postMessage(
+          {
+            perfetto: {
+              buffer: traceData,
+              title: `Trace ${taskId}`,
+            },
+          },
+          "https://ui.perfetto.dev",
+        );
+      }
+    };
+    window.addEventListener("message", onMessage);
+
+    // Ping Perfetto until it responds
+    const pingInterval = setInterval(() => {
+      perfettoWindow.postMessage("PING", "https://ui.perfetto.dev");
+    }, 50);
+
+    // Clean up after 10 seconds
+    setTimeout(() => {
+      clearInterval(pingInterval);
+      window.removeEventListener("message", onMessage);
+    }, 10000);
+  }
+
+  async function refresh() {
+    loading = true;
+    error = null;
+    try {
+      const response = await listTraces();
+      traces = response.traces;
+    } catch (e) {
+      error = e instanceof Error ? e.message : "Failed to load traces";
+    } finally {
+      loading = false;
+    }
+  }
+
+  onMount(() => {
+    refresh();
+  });
+</script>
+
+<div class="min-h-screen bg-exo-dark-gray text-white">
+  <HeaderNav showHome={true} />
+  <div class="max-w-7xl mx-auto px-4 lg:px-8 py-6 space-y-6">
+    <div class="flex items-center justify-between gap-4 flex-wrap">
+      <div>
+        <h1
+          class="text-2xl font-mono tracking-[0.2em] uppercase text-exo-yellow"
+        >
+          Traces
+        </h1>
+      </div>
+      <div class="flex items-center gap-3">
+        <button
+          type="button"
+          class="text-xs font-mono text-exo-light-gray hover:text-exo-yellow transition-colors uppercase border border-exo-medium-gray/40 px-2 py-1 rounded"
+          onclick={refresh}
+          disabled={loading}
+        >
+          Refresh
+        </button>
+      </div>
+    </div>
+
+    {#if loading}
+      <div
+        class="rounded border border-exo-medium-gray/30 bg-exo-black/30 p-6 text-center text-exo-light-gray"
+      >
+        <div class="text-sm">Loading traces...</div>
+      </div>
+    {:else if error}
+      <div
+        class="rounded border border-red-500/30 bg-red-500/10 p-6 text-center text-red-400"
+      >
+        <div class="text-sm">{error}</div>
+      </div>
+    {:else if traces.length === 0}
+      <div
+        class="rounded border border-exo-medium-gray/30 bg-exo-black/30 p-6 text-center text-exo-light-gray space-y-2"
+      >
+        <div class="text-sm">No traces found.</div>
+        <div class="text-xs text-exo-light-gray/70">
+          Run exo with EXO_TRACING_ENABLED=1 to collect traces.
+        </div>
+      </div>
+    {:else}
+      <div class="space-y-3">
+        {#each traces as trace}
+          <div
+            class="rounded border border-exo-medium-gray/30 bg-exo-black/30 p-4 flex items-center justify-between gap-4"
+          >
+            <div class="min-w-0 flex-1">
+              <a
+                href="#/traces/{trace.taskId}"
+                class="text-sm font-mono text-white hover:text-exo-yellow transition-colors truncate block"
+              >
+                {trace.taskId}
+              </a>
+              <div class="text-xs text-exo-light-gray font-mono mt-1">
+                {formatDate(trace.createdAt)} &bull; {formatBytes(
+                  trace.fileSize,
+                )}
+              </div>
+            </div>
+            <div class="flex items-center gap-2 shrink-0">
+              <a
+                href="#/traces/{trace.taskId}"
+                class="text-xs font-mono text-exo-light-gray hover:text-exo-yellow transition-colors uppercase border border-exo-medium-gray/40 px-2 py-1 rounded"
+              >
+                View Stats
+              </a>
+              <button
+                type="button"
+                class="text-xs font-mono text-exo-light-gray hover:text-exo-yellow transition-colors uppercase border border-exo-medium-gray/40 px-2 py-1 rounded"
+                onclick={() => downloadTrace(trace.taskId)}
+              >
+                Download
+              </button>
+              <button
+                type="button"
+                class="text-xs font-mono text-exo-dark-gray bg-exo-yellow hover:bg-exo-yellow/90 transition-colors uppercase px-2 py-1 rounded font-semibold"
+                onclick={() => openInPerfetto(trace.taskId)}
+              >
+                View Trace
+              </button>
+            </div>
+          </div>
+        {/each}
+      </div>
+    {/if}
+  </div>
+</div>
diff --git a/dashboard/src/routes/traces/[taskId]/+page.svelte b/dashboard/src/routes/traces/[taskId]/+page.svelte
new file mode 100644
index 00000000..f370c8fa
--- /dev/null
+++ b/dashboard/src/routes/traces/[taskId]/+page.svelte
@@ -0,0 +1,367 @@
+<script lang="ts">
+  import { page } from "$app/stores";
+  import { onMount } from "svelte";
+  import {
+    fetchTraceStats,
+    getTraceRawUrl,
+    type TraceStatsResponse,
+    type TraceCategoryStats,
+  } from "$lib/stores/app.svelte";
+  import HeaderNav from "$lib/components/HeaderNav.svelte";
+
+  const taskId = $derived($page.params.taskId);
+
+  let stats = $state<TraceStatsResponse | null>(null);
+  let loading = $state(true);
+  let error = $state<string | null>(null);
+
+  function formatDuration(us: number): string {
+    if (us < 1000) return `${us.toFixed(0)}us`;
+    if (us < 1_000_000) return `${(us / 1000).toFixed(2)}ms`;
+    return `${(us / 1_000_000).toFixed(2)}s`;
+  }
+
+  function formatPercentage(part: number, total: number): string {
+    if (total === 0) return "0.0%";
+    return `${((part / total) * 100).toFixed(1)}%`;
+  }
+
+  // Parse hierarchical categories like "sync/compute" into phases
+  type PhaseData = {
+    name: string;
+    subcategories: { name: string; stats: TraceCategoryStats }[];
+    totalUs: number; // From outer span (e.g., "sync" category)
+    stepCount: number; // Count of outer span events
+  };
+
+  function parsePhases(
+    byCategory: Record<string, TraceCategoryStats>,
+  ): PhaseData[] {
+    const phases = new Map<
+      string,
+      {
+        subcats: Map<string, TraceCategoryStats>;
+        outerStats: TraceCategoryStats | null;
+      }
+    >();
+
+    for (const [category, catStats] of Object.entries(byCategory)) {
+      if (category.includes("/")) {
+        const [phase, subcat] = category.split("/", 2);
+        if (!phases.has(phase)) {
+          phases.set(phase, { subcats: new Map(), outerStats: null });
+        }
+        phases.get(phase)!.subcats.set(subcat, catStats);
+      } else {
+        // Outer span - this IS the phase total
+        if (!phases.has(category)) {
+          phases.set(category, { subcats: new Map(), outerStats: null });
+        }
+        phases.get(category)!.outerStats = catStats;
+      }
+    }
+
+    return Array.from(phases.entries())
+      .filter(([_, data]) => data.outerStats !== null) // Only phases with outer spans
+      .map(([name, data]) => ({
+        name,
+        subcategories: Array.from(data.subcats.entries())
+          .map(([subName, subStats]) => ({ name: subName, stats: subStats }))
+          .sort((a, b) => b.stats.totalUs - a.stats.totalUs),
+        totalUs: data.outerStats!.totalUs, // Outer span total
+        stepCount: data.outerStats!.count, // Number of steps
+      }))
+      .sort((a, b) => b.totalUs - a.totalUs);
+  }
+
+  async function downloadTrace() {
+    if (!taskId) return;
+    const response = await fetch(getTraceRawUrl(taskId));
+    const blob = await response.blob();
+    const url = URL.createObjectURL(blob);
+    const a = document.createElement("a");
+    a.href = url;
+    a.download = `trace_${taskId}.json`;
+    a.click();
+    URL.revokeObjectURL(url);
+  }
+
+  async function openInPerfetto() {
+    if (!taskId) return;
+
+    // Fetch trace data from our local API
+    const response = await fetch(getTraceRawUrl(taskId));
+    const traceData = await response.arrayBuffer();
+
+    // Open Perfetto UI
+    const perfettoWindow = window.open("https://ui.perfetto.dev");
+    if (!perfettoWindow) {
+      alert("Failed to open Perfetto. Please allow popups.");
+      return;
+    }
+
+    // Wait for Perfetto to be ready, then send trace via postMessage
+    const onMessage = (e: MessageEvent) => {
+      if (e.data === "PONG") {
+        window.removeEventListener("message", onMessage);
+        perfettoWindow.postMessage(
+          {
+            perfetto: {
+              buffer: traceData,
+              title: `Trace ${taskId}`,
+            },
+          },
+          "https://ui.perfetto.dev",
+        );
+      }
+    };
+    window.addEventListener("message", onMessage);
+
+    // Ping Perfetto until it responds
+    const pingInterval = setInterval(() => {
+      perfettoWindow.postMessage("PING", "https://ui.perfetto.dev");
+    }, 50);
+
+    // Clean up after 10 seconds
+    setTimeout(() => {
+      clearInterval(pingInterval);
+      window.removeEventListener("message", onMessage);
+    }, 10000);
+  }
+
+  onMount(async () => {
+    if (!taskId) {
+      error = "No task ID provided";
+      loading = false;
+      return;
+    }
+
+    try {
+      stats = await fetchTraceStats(taskId);
+    } catch (e) {
+      error = e instanceof Error ? e.message : "Failed to load trace";
+    } finally {
+      loading = false;
+    }
+  });
+
+  const phases = $derived(stats ? parsePhases(stats.byCategory) : []);
+  const sortedRanks = $derived(
+    stats
+      ? Object.keys(stats.byRank)
+          .map(Number)
+          .sort((a, b) => a - b)
+      : [],
+  );
+  const nodeCount = $derived(sortedRanks.length || 1);
+</script>
+
+<div class="min-h-screen bg-exo-dark-gray text-white">
+  <HeaderNav showHome={true} />
+  <div class="max-w-7xl mx-auto px-4 lg:px-8 py-6 space-y-6">
+    <div class="flex items-center justify-between gap-4 flex-wrap">
+      <div>
+        <h1
+          class="text-2xl font-mono tracking-[0.2em] uppercase text-exo-yellow"
+        >
+          Trace
+        </h1>
+        <p class="text-sm text-exo-light-gray font-mono truncate max-w-lg">
+          {taskId}
+        </p>
+      </div>
+      <div class="flex items-center gap-3">
+        <a
+          href="#/traces"
+          class="text-xs font-mono text-exo-light-gray hover:text-exo-yellow transition-colors uppercase border border-exo-medium-gray/40 px-3 py-1.5 rounded"
+        >
+          All Traces
+        </a>
+        <button
+          type="button"
+          class="text-xs font-mono text-exo-light-gray hover:text-exo-yellow transition-colors uppercase border border-exo-medium-gray/40 px-3 py-1.5 rounded"
+          onclick={downloadTrace}
+          disabled={loading || !!error}
+        >
+          Download
+        </button>
+        <button
+          type="button"
+          class="text-xs font-mono text-exo-dark-gray bg-exo-yellow hover:bg-exo-yellow/90 transition-colors uppercase px-3 py-1.5 rounded font-semibold"
+          onclick={openInPerfetto}
+          disabled={loading || !!error}
+        >
+          View Trace
+        </button>
+      </div>
+    </div>
+
+    {#if loading}
+      <div
+        class="rounded border border-exo-medium-gray/30 bg-exo-black/30 p-6 text-center text-exo-light-gray"
+      >
+        <div class="text-sm">Loading trace data...</div>
+      </div>
+    {:else if error}
+      <div
+        class="rounded border border-red-500/30 bg-red-500/10 p-6 text-center text-red-400"
+      >
+        <div class="text-sm">{error}</div>
+      </div>
+    {:else if stats}
+      <!-- Wall Time Summary -->
+      <div
+        class="rounded border border-exo-medium-gray/30 bg-exo-black/30 p-4 space-y-2"
+      >
+        <h2
+          class="text-sm font-mono uppercase tracking-wider text-exo-light-gray"
+        >
+          Summary
+        </h2>
+        <div class="text-3xl font-mono text-exo-yellow">
+          {formatDuration(stats.totalWallTimeUs)}
+        </div>
+        <div class="text-xs text-exo-light-gray">Total wall time</div>
+      </div>
+
+      <!-- By Phase -->
+      {#if phases.length > 0}
+        <div
+          class="rounded border border-exo-medium-gray/30 bg-exo-black/30 p-4 space-y-4"
+        >
+          <h2
+            class="text-sm font-mono uppercase tracking-wider text-exo-light-gray"
+          >
+            By Phase <span class="text-exo-light-gray/50">(avg per node)</span>
+          </h2>
+          <div class="space-y-4">
+            {#each phases as phase}
+              {@const normalizedTotal = phase.totalUs / nodeCount}
+              {@const normalizedStepCount = phase.stepCount / nodeCount}
+              <div class="space-y-2">
+                <div class="flex items-center justify-between">
+                  <span class="text-sm font-mono text-white">{phase.name}</span>
+                  <span class="text-sm font-mono">
+                    <span class="text-exo-yellow"
+                      >{formatDuration(normalizedTotal)}</span
+                    >
+                    <span class="text-exo-light-gray ml-2">
+                      ({normalizedStepCount} steps, {formatDuration(
+                        normalizedTotal / normalizedStepCount,
+                      )}/step)
+                    </span>
+                  </span>
+                </div>
+                {#if phase.subcategories.length > 0}
+                  <div class="pl-4 space-y-1.5">
+                    {#each phase.subcategories as subcat}
+                      {@const normalizedSubcat =
+                        subcat.stats.totalUs / nodeCount}
+                      {@const pct = formatPercentage(
+                        normalizedSubcat,
+                        normalizedTotal,
+                      )}
+                      {@const perStep = normalizedSubcat / normalizedStepCount}
+                      <div
+                        class="flex items-center justify-between text-xs font-mono"
+                      >
+                        <span class="text-exo-light-gray">{subcat.name}</span>
+                        <span class="text-white">
+                          {formatDuration(normalizedSubcat)}
+                          <span class="text-exo-light-gray ml-2">({pct})</span>
+                          <span class="text-exo-light-gray/60 ml-2"
+                            >{formatDuration(perStep)}/step</span
+                          >
+                        </span>
+                      </div>
+                      <!-- Progress bar -->
+                      <div
+                        class="relative h-1.5 bg-exo-black/60 rounded-sm overflow-hidden"
+                      >
+                        <div
+                          class="absolute inset-y-0 left-0 bg-gradient-to-r from-exo-yellow to-exo-yellow/70 transition-all duration-300"
+                          style="width: {pct}"
+                        ></div>
+                      </div>
+                    {/each}
+                  </div>
+                {/if}
+              </div>
+            {/each}
+          </div>
+        </div>
+      {/if}
+
+      <!-- By Rank -->
+      {#if sortedRanks.length > 0}
+        <div
+          class="rounded border border-exo-medium-gray/30 bg-exo-black/30 p-4 space-y-4"
+        >
+          <h2
+            class="text-sm font-mono uppercase tracking-wider text-exo-light-gray"
+          >
+            By Rank
+          </h2>
+          <div class="grid grid-cols-1 md:grid-cols-2 lg:grid-cols-3 gap-4">
+            {#each sortedRanks as rank}
+              {@const rankStats = stats.byRank[rank]}
+              {@const rankPhases = parsePhases(rankStats.byCategory)}
+              <div
+                class="rounded border border-exo-medium-gray/20 bg-exo-dark-gray/60 p-3 space-y-3"
+              >
+                <div class="text-sm font-mono text-exo-yellow">
+                  Rank {rank}
+                </div>
+                <div class="space-y-2">
+                  {#each rankPhases as phase}
+                    <div class="space-y-1">
+                      <div class="flex items-center justify-between text-xs">
+                        <span class="font-mono text-exo-light-gray"
+                          >{phase.name}</span
+                        >
+                        <span class="font-mono text-white">
+                          {formatDuration(phase.totalUs)}
+                          <span class="text-exo-light-gray/50 ml-1">
+                            ({phase.stepCount}x)
+                          </span>
+                        </span>
+                      </div>
+                      {#if phase.subcategories.length > 0}
+                        <div class="pl-2 space-y-0.5">
+                          {#each phase.subcategories as subcat}
+                            {@const pct = formatPercentage(
+                              subcat.stats.totalUs,
+                              phase.totalUs,
+                            )}
+                            {@const perStep =
+                              subcat.stats.totalUs / phase.stepCount}
+                            <div
+                              class="flex items-center justify-between text-[10px] font-mono"
+                            >
+                              <span class="text-exo-light-gray/70"
+                                >{subcat.name}</span
+                              >
+                              <span class="text-exo-light-gray">
+                                {formatDuration(subcat.stats.totalUs)}
+                                <span class="text-exo-light-gray/50"
+                                  >({pct})</span
+                                >
+                                <span class="text-exo-light-gray/30 ml-1"
+                                  >{formatDuration(perStep)}/step</span
+                                >
+                              </span>
+                            </div>
+                          {/each}
+                        </div>
+                      {/if}
+                    </div>
+                  {/each}
+                </div>
+              </div>
+            {/each}
+          </div>
+        </div>
+      {/if}
+    {/if}
+  </div>
+</div>
diff --git a/src/exo/master/api.py b/src/exo/master/api.py
index 908716d8..4855e62a 100644
--- a/src/exo/master/api.py
+++ b/src/exo/master/api.py
@@ -3,7 +3,9 @@ import contextlib
 import json
 import time
 from collections.abc import AsyncGenerator, Awaitable, Callable
+from datetime import datetime, timezone
 from http import HTTPStatus
+from pathlib import Path
 from typing import Annotated, Literal, cast
 from uuid import uuid4
 
@@ -40,6 +42,7 @@ from exo.shared.apply import apply
 from exo.shared.constants import (
     EXO_IMAGE_CACHE_DIR,
     EXO_MAX_CHUNK_SIZE,
+    EXO_TRACING_CACHE_DIR,
 )
 from exo.shared.election import ElectionMessage
 from exo.shared.logging import InterceptLogger
@@ -48,6 +51,7 @@ from exo.shared.models.model_cards import (
     ModelCard,
     ModelId,
 )
+from exo.shared.tracing import TraceEvent, compute_stats, export_trace, load_trace_file
 from exo.shared.types.api import (
     AdvancedImageParams,
     BenchChatCompletionRequest,
@@ -81,6 +85,13 @@ from exo.shared.types.api import (
     StartDownloadParams,
     StartDownloadResponse,
     ToolCall,
+    TraceCategoryStats,
+    TraceEventResponse,
+    TraceListItem,
+    TraceListResponse,
+    TraceRankStats,
+    TraceResponse,
+    TraceStatsResponse,
 )
 from exo.shared.types.chunks import (
     ErrorChunk,
@@ -115,6 +126,7 @@ from exo.shared.types.events import (
     Event,
     ForwarderEvent,
     IndexedEvent,
+    TracesMerged,
 )
 from exo.shared.types.memory import Memory
 from exo.shared.types.openai_responses import (
@@ -275,6 +287,10 @@ class API:
         self.app.get("/events")(lambda: self._event_log)
         self.app.post("/download/start")(self.start_download)
         self.app.delete("/download/{node_id}/{model_id:path}")(self.delete_download)
+        self.app.get("/v1/traces")(self.list_traces)
+        self.app.get("/v1/traces/{task_id}")(self.get_trace)
+        self.app.get("/v1/traces/{task_id}/stats")(self.get_trace_stats)
+        self.app.get("/v1/traces/{task_id}/raw")(self.get_trace_raw)
 
     async def place_instance(self, payload: PlaceInstanceParams):
         command = PlaceInstance(
@@ -1279,6 +1295,24 @@ class API:
                             except BrokenResourceError:
                                 self._text_generation_queues.pop(event.command_id, None)
 
+                    if isinstance(event, TracesMerged):
+                        self._save_merged_trace(event)
+
+    def _save_merged_trace(self, event: TracesMerged) -> None:
+        traces = [
+            TraceEvent(
+                name=t.name,
+                start_us=t.start_us,
+                duration_us=t.duration_us,
+                rank=t.rank,
+                category=t.category,
+            )
+            for t in event.traces
+        ]
+        output_path = EXO_TRACING_CACHE_DIR / f"trace_{event.task_id}.json"
+        export_trace(traces, output_path)
+        logger.debug(f"Saved merged trace to {output_path}")
+
     async def _pause_on_new_election(self):
         with self.election_receiver as ems:
             async for message in ems:
@@ -1325,3 +1359,103 @@ class API:
         )
         await self._send_download(command)
         return DeleteDownloadResponse(command_id=command.command_id)
+
+    def _get_trace_path(self, task_id: str) -> Path:
+        return EXO_TRACING_CACHE_DIR / f"trace_{task_id}.json"
+
+    async def list_traces(self) -> TraceListResponse:
+        traces: list[TraceListItem] = []
+
+        for trace_file in sorted(
+            EXO_TRACING_CACHE_DIR.glob("trace_*.json"),
+            key=lambda p: p.stat().st_mtime,
+            reverse=True,
+        ):
+            # Extract task_id from filename (trace_{task_id}.json)
+            task_id = trace_file.stem.removeprefix("trace_")
+            stat = trace_file.stat()
+            created_at = datetime.fromtimestamp(
+                stat.st_mtime, tz=timezone.utc
+            ).isoformat()
+            traces.append(
+                TraceListItem(
+                    task_id=task_id,
+                    created_at=created_at,
+                    file_size=stat.st_size,
+                )
+            )
+
+        return TraceListResponse(traces=traces)
+
+    async def get_trace(self, task_id: str) -> TraceResponse:
+        trace_path = self._get_trace_path(task_id)
+
+        if not trace_path.exists():
+            raise HTTPException(status_code=404, detail=f"Trace not found: {task_id}")
+
+        trace_events = load_trace_file(trace_path)
+
+        return TraceResponse(
+            task_id=task_id,
+            traces=[
+                TraceEventResponse(
+                    name=event.name,
+                    start_us=event.start_us,
+                    duration_us=event.duration_us,
+                    rank=event.rank,
+                    category=event.category,
+                )
+                for event in trace_events
+            ],
+        )
+
+    async def get_trace_stats(self, task_id: str) -> TraceStatsResponse:
+        trace_path = self._get_trace_path(task_id)
+
+        if not trace_path.exists():
+            raise HTTPException(status_code=404, detail=f"Trace not found: {task_id}")
+
+        trace_events = load_trace_file(trace_path)
+        stats = compute_stats(trace_events)
+
+        return TraceStatsResponse(
+            task_id=task_id,
+            total_wall_time_us=stats.total_wall_time_us,
+            by_category={
+                category: TraceCategoryStats(
+                    total_us=cat_stats.total_us,
+                    count=cat_stats.count,
+                    min_us=cat_stats.min_us,
+                    max_us=cat_stats.max_us,
+                    avg_us=cat_stats.avg_us,
+                )
+                for category, cat_stats in stats.by_category.items()
+            },
+            by_rank={
+                rank: TraceRankStats(
+                    by_category={
+                        category: TraceCategoryStats(
+                            total_us=cat_stats.total_us,
+                            count=cat_stats.count,
+                            min_us=cat_stats.min_us,
+                            max_us=cat_stats.max_us,
+                            avg_us=cat_stats.avg_us,
+                        )
+                        for category, cat_stats in rank_stats.items()
+                    }
+                )
+                for rank, rank_stats in stats.by_rank.items()
+            },
+        )
+
+    async def get_trace_raw(self, task_id: str) -> FileResponse:
+        trace_path = self._get_trace_path(task_id)
+
+        if not trace_path.exists():
+            raise HTTPException(status_code=404, detail=f"Trace not found: {task_id}")
+
+        return FileResponse(
+            path=trace_path,
+            media_type="application/json",
+            filename=f"trace_{task_id}.json",
+        )
diff --git a/src/exo/master/main.py b/src/exo/master/main.py
index d99a96c9..a3041143 100644
--- a/src/exo/master/main.py
+++ b/src/exo/master/main.py
@@ -11,6 +11,7 @@ from exo.master.placement import (
     place_instance,
 )
 from exo.shared.apply import apply
+from exo.shared.constants import EXO_TRACING_ENABLED
 from exo.shared.types.commands import (
     CreateInstance,
     DeleteInstance,
@@ -35,6 +36,9 @@ from exo.shared.types.events import (
     NodeTimedOut,
     TaskCreated,
     TaskDeleted,
+    TraceEventData,
+    TracesCollected,
+    TracesMerged,
 )
 from exo.shared.types.state import State
 from exo.shared.types.tasks import (
@@ -86,6 +90,8 @@ class Master:
         self._multi_buffer = MultiSourceBuffer[NodeId, Event]()
         # TODO: not have this
         self._event_log: list[Event] = []
+        self._pending_traces: dict[TaskId, dict[int, list[TraceEventData]]] = {}
+        self._expected_ranks: dict[TaskId, set[int]] = {}
 
     async def run(self):
         logger.info("Starting Master")
@@ -187,13 +193,14 @@ class Master:
                             )
 
                             task_id = TaskId()
+                            selected_instance_id = available_instance_ids[0]
                             generated_events.append(
                                 TaskCreated(
                                     task_id=task_id,
                                     task=ImageGenerationTask(
                                         task_id=task_id,
                                         command_id=command.command_id,
-                                        instance_id=available_instance_ids[0],
+                                        instance_id=selected_instance_id,
                                         task_status=TaskStatus.Pending,
                                         task_params=command.task_params,
                                     ),
@@ -201,6 +208,17 @@ class Master:
                             )
 
                             self.command_task_mapping[command.command_id] = task_id
+
+                            if EXO_TRACING_ENABLED:
+                                selected_instance = self.state.instances.get(
+                                    selected_instance_id
+                                )
+                                if selected_instance:
+                                    ranks = set(
+                                        shard.device_rank
+                                        for shard in selected_instance.shard_assignments.runner_to_shard.values()
+                                    )
+                                    self._expected_ranks[task_id] = ranks
                         case ImageEdits():
                             for instance in self.state.instances.values():
                                 if (
@@ -229,13 +247,14 @@ class Master:
                             )
 
                             task_id = TaskId()
+                            selected_instance_id = available_instance_ids[0]
                             generated_events.append(
                                 TaskCreated(
                                     task_id=task_id,
                                     task=ImageEditsTask(
                                         task_id=task_id,
                                         command_id=command.command_id,
-                                        instance_id=available_instance_ids[0],
+                                        instance_id=selected_instance_id,
                                         task_status=TaskStatus.Pending,
                                         task_params=command.task_params,
                                     ),
@@ -243,6 +262,17 @@ class Master:
                             )
 
                             self.command_task_mapping[command.command_id] = task_id
+
+                            if EXO_TRACING_ENABLED:
+                                selected_instance = self.state.instances.get(
+                                    selected_instance_id
+                                )
+                                if selected_instance:
+                                    ranks = set(
+                                        shard.device_rank
+                                        for shard in selected_instance.shard_assignments.runner_to_shard.values()
+                                    )
+                                    self._expected_ranks[task_id] = ranks
                         case DeleteInstance():
                             placement = delete_instance(command, self.state.instances)
                             transition_events = get_transition_events(
@@ -335,6 +365,10 @@ class Master:
                     local_event.origin,
                 )
                 for event in self._multi_buffer.drain():
+                    if isinstance(event, TracesCollected):
+                        await self._handle_traces_collected(event)
+                        continue
+
                     logger.debug(f"Master indexing event: {str(event)[:100]}")
                     indexed = IndexedEvent(event=event, idx=len(self._event_log))
                     self.state = apply(self.state, indexed)
@@ -373,3 +407,29 @@ class Master:
                 event=event.event,
             )
         )
+
+    async def _handle_traces_collected(self, event: TracesCollected) -> None:
+        task_id = event.task_id
+        if task_id not in self._pending_traces:
+            self._pending_traces[task_id] = {}
+        self._pending_traces[task_id][event.rank] = event.traces
+
+        if (
+            task_id in self._expected_ranks
+            and set(self._pending_traces[task_id].keys())
+            >= self._expected_ranks[task_id]
+        ):
+            await self._merge_and_save_traces(task_id)
+
+    async def _merge_and_save_traces(self, task_id: TaskId) -> None:
+        all_trace_data: list[TraceEventData] = []
+        for trace_data in self._pending_traces[task_id].values():
+            all_trace_data.extend(trace_data)
+
+        await self.event_sender.send(
+            TracesMerged(task_id=task_id, traces=all_trace_data)
+        )
+
+        del self._pending_traces[task_id]
+        if task_id in self._expected_ranks:
+            del self._expected_ranks[task_id]
diff --git a/src/exo/shared/apply.py b/src/exo/shared/apply.py
index 675b9647..8c2304fe 100644
--- a/src/exo/shared/apply.py
+++ b/src/exo/shared/apply.py
@@ -25,6 +25,8 @@ from exo.shared.types.events import (
     TestEvent,
     TopologyEdgeCreated,
     TopologyEdgeDeleted,
+    TracesCollected,
+    TracesMerged,
 )
 from exo.shared.types.profiling import (
     NodeIdentity,
@@ -55,7 +57,12 @@ def event_apply(event: Event, state: State) -> State:
     """Apply an event to state."""
     match event:
         case (
-            TestEvent() | ChunkGenerated() | TaskAcknowledged() | InputChunkReceived()
+            TestEvent()
+            | ChunkGenerated()
+            | TaskAcknowledged()
+            | InputChunkReceived()
+            | TracesCollected()
+            | TracesMerged()
         ):  # Pass-through events that don't modify state
             return state
         case InstanceCreated():
diff --git a/src/exo/shared/constants.py b/src/exo/shared/constants.py
index dfd610b4..ad503e90 100644
--- a/src/exo/shared/constants.py
+++ b/src/exo/shared/constants.py
@@ -49,7 +49,10 @@ LIBP2P_COMMANDS_TOPIC = "commands"
 EXO_MAX_CHUNK_SIZE = 512 * 1024
 
 EXO_IMAGE_CACHE_DIR = EXO_CACHE_HOME / "images"
+EXO_TRACING_CACHE_DIR = EXO_CACHE_HOME / "traces"
 
 EXO_ENABLE_IMAGE_MODELS = (
     os.getenv("EXO_ENABLE_IMAGE_MODELS", "false").lower() == "true"
 )
+
+EXO_TRACING_ENABLED = os.getenv("EXO_TRACING_ENABLED", "false").lower() == "true"
diff --git a/src/exo/shared/tracing.py b/src/exo/shared/tracing.py
new file mode 100644
index 00000000..a353d923
--- /dev/null
+++ b/src/exo/shared/tracing.py
@@ -0,0 +1,238 @@
+from __future__ import annotations
+
+import json
+import time
+from collections import defaultdict
+from collections.abc import Generator
+from contextlib import contextmanager
+from contextvars import ContextVar
+from dataclasses import dataclass, field
+from pathlib import Path
+from typing import cast, final
+
+from exo.shared.constants import EXO_TRACING_ENABLED
+from exo.worker.runner.bootstrap import logger
+
+# Context variable to track the current trace category for hierarchical nesting
+_current_category: ContextVar[str | None] = ContextVar("current_category", default=None)
+
+
+@final
+@dataclass(frozen=True)
+class TraceEvent:
+    name: str
+    start_us: int
+    duration_us: int
+    rank: int
+    category: str
+
+
+@final
+@dataclass
+class CategoryStats:
+    total_us: int = 0
+    count: int = 0
+    min_us: int = 0
+    max_us: int = 0
+
+    def add(self, duration_us: int) -> None:
+        if self.count == 0:
+            self.min_us = duration_us
+            self.max_us = duration_us
+        else:
+            self.min_us = min(self.min_us, duration_us)
+            self.max_us = max(self.max_us, duration_us)
+        self.total_us += duration_us
+        self.count += 1
+
+    @property
+    def avg_us(self) -> float:
+        return self.total_us / self.count if self.count > 0 else 0.0
+
+
+@final
+@dataclass
+class TraceStats:
+    total_wall_time_us: int = 0
+    by_category: dict[str, CategoryStats] = field(default_factory=dict)
+    by_rank: dict[int, dict[str, CategoryStats]] = field(default_factory=dict)
+
+
+# Global trace buffer - each rank accumulates traces here
+_trace_buffer: list[TraceEvent] = []
+
+
+def _record_span(
+    name: str, start_us: int, duration_us: int, rank: int, category: str
+) -> None:
+    _trace_buffer.append(
+        TraceEvent(
+            name=name,
+            start_us=start_us,
+            duration_us=duration_us,
+            rank=rank,
+            category=category,
+        )
+    )
+
+
+@contextmanager
+def trace(
+    name: str,
+    rank: int,
+    category: str = "compute",
+) -> Generator[None, None, None]:
+    """Context manager to trace any operation.
+
+    Nested traces automatically inherit the parent category, creating hierarchical
+    categories like "sync/compute" or "async/comms".
+
+    Args:
+        name: Name of the operation (e.g., "recv 0", "send 1", "joint_blocks")
+        rank: This rank's ID
+        category: Category for grouping in trace viewer ("comm", "compute", "step")
+
+    Example:
+        with trace(f"sync {t}", rank, "sync"):
+            with trace("joint_blocks", rank, "compute"):
+                # Recorded with category "sync/compute"
+                hidden_states = some_computation(...)
+    """
+    if not EXO_TRACING_ENABLED:
+        yield
+        return
+
+    # Combine with parent category if nested
+    parent = _current_category.get()
+    full_category = f"{parent}/{category}" if parent else category
+
+    # Set as current for nested traces
+    token = _current_category.set(full_category)
+
+    try:
+        start_us = int(time.time() * 1_000_000)
+        start_perf = time.perf_counter()
+        yield
+        duration_us = int((time.perf_counter() - start_perf) * 1_000_000)
+        _record_span(name, start_us, duration_us, rank, full_category)
+    finally:
+        _current_category.reset(token)
+
+
+def get_trace_buffer() -> list[TraceEvent]:
+    return list(_trace_buffer)
+
+
+def clear_trace_buffer() -> None:
+    _trace_buffer.clear()
+
+
+def export_trace(traces: list[TraceEvent], output_path: Path) -> None:
+    trace_events: list[dict[str, object]] = []
+
+    for event in traces:
+        # Chrome trace format uses "X" for complete events (with duration)
+        chrome_event: dict[str, object] = {
+            "name": event.name,
+            "cat": event.category,
+            "ph": "X",
+            "ts": event.start_us,
+            "dur": event.duration_us,
+            "pid": 0,
+            "tid": event.rank,
+            "args": {"rank": event.rank},
+        }
+        trace_events.append(chrome_event)
+
+    ranks_seen = set(t.rank for t in traces)
+    for rank in ranks_seen:
+        trace_events.append(
+            {
+                "name": "thread_name",
+                "ph": "M",  # Metadata event
+                "pid": 0,
+                "tid": rank,
+                "args": {"name": f"Rank {rank}"},
+            }
+        )
+
+    chrome_trace = {"traceEvents": trace_events}
+
+    try:
+        output_path.parent.mkdir(parents=True, exist_ok=True)
+        with open(output_path, "w") as f:
+            json.dump(chrome_trace, f, indent=2)
+    except OSError as e:
+        logger.warning("Failed to export trace to %s: %s", output_path, e)
+
+
+def load_trace_file(path: Path) -> list[TraceEvent]:
+    with open(path) as f:
+        data = cast(dict[str, list[dict[str, object]]], json.load(f))
+
+    events = data.get("traceEvents", [])
+    traces: list[TraceEvent] = []
+
+    for event in events:
+        # Skip metadata events
+        if event.get("ph") == "M":
+            continue
+
+        name = str(event.get("name", ""))
+        category = str(event.get("cat", ""))
+        ts_value = event.get("ts", 0)
+        dur_value = event.get("dur", 0)
+        tid_value = event.get("tid", 0)
+        start_us = int(ts_value) if isinstance(ts_value, (int, float, str)) else 0
+        duration_us = int(dur_value) if isinstance(dur_value, (int, float, str)) else 0
+
+        # Get rank from tid or args
+        rank = int(tid_value) if isinstance(tid_value, (int, float, str)) else 0
+        args = event.get("args")
+        if isinstance(args, dict):
+            args_dict = cast(dict[str, object], args)
+            rank_from_args = args_dict.get("rank")
+            if isinstance(rank_from_args, (int, float, str)):
+                rank = int(rank_from_args)
+
+        traces.append(
+            TraceEvent(
+                name=name,
+                start_us=start_us,
+                duration_us=duration_us,
+                rank=rank,
+                category=category,
+            )
+        )
+
+    return traces
+
+
+def compute_stats(traces: list[TraceEvent]) -> TraceStats:
+    stats = TraceStats()
+
+    if not traces:
+        return stats
+
+    # Calculate wall time from earliest start to latest end
+    min_start = min(t.start_us for t in traces)
+    max_end = max(t.start_us + t.duration_us for t in traces)
+    stats.total_wall_time_us = max_end - min_start
+
+    # Initialize nested dicts
+    by_category: dict[str, CategoryStats] = defaultdict(CategoryStats)
+    by_rank: dict[int, dict[str, CategoryStats]] = defaultdict(
+        lambda: defaultdict(CategoryStats)
+    )
+
+    for event in traces:
+        # By category
+        by_category[event.category].add(event.duration_us)
+
+        # By rank and category
+        by_rank[event.rank][event.category].add(event.duration_us)
+
+    stats.by_category = dict(by_category)
+    stats.by_rank = {k: dict(v) for k, v in by_rank.items()}
+
+    return stats
diff --git a/src/exo/shared/types/api.py b/src/exo/shared/types/api.py
index 3f665179..40dbb288 100644
--- a/src/exo/shared/types/api.py
+++ b/src/exo/shared/types/api.py
@@ -350,3 +350,45 @@ class StartDownloadResponse(CamelCaseModel):
 
 class DeleteDownloadResponse(CamelCaseModel):
     command_id: CommandId
+
+
+class TraceEventResponse(CamelCaseModel):
+    name: str
+    start_us: int
+    duration_us: int
+    rank: int
+    category: str
+
+
+class TraceResponse(CamelCaseModel):
+    task_id: str
+    traces: list[TraceEventResponse]
+
+
+class TraceCategoryStats(CamelCaseModel):
+    total_us: int
+    count: int
+    min_us: int
+    max_us: int
+    avg_us: float
+
+
+class TraceRankStats(CamelCaseModel):
+    by_category: dict[str, TraceCategoryStats]
+
+
+class TraceStatsResponse(CamelCaseModel):
+    task_id: str
+    total_wall_time_us: int
+    by_category: dict[str, TraceCategoryStats]
+    by_rank: dict[int, TraceRankStats]
+
+
+class TraceListItem(CamelCaseModel):
+    task_id: str
+    created_at: str
+    file_size: int
+
+
+class TraceListResponse(CamelCaseModel):
+    traces: list[TraceListItem]
diff --git a/src/exo/shared/types/events.py b/src/exo/shared/types/events.py
index 647e510f..5cf93d0c 100644
--- a/src/exo/shared/types/events.py
+++ b/src/exo/shared/types/events.py
@@ -1,4 +1,5 @@
 from datetime import datetime
+from typing import final
 
 from pydantic import Field
 
@@ -10,7 +11,7 @@ from exo.shared.types.worker.downloads import DownloadProgress
 from exo.shared.types.worker.instances import Instance, InstanceId
 from exo.shared.types.worker.runners import RunnerId, RunnerStatus
 from exo.utils.info_gatherer.info_gatherer import GatheredInfo
-from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
+from exo.utils.pydantic_ext import CamelCaseModel, FrozenModel, TaggedModel
 
 
 class EventId(Id):
@@ -109,6 +110,28 @@ class TopologyEdgeDeleted(BaseEvent):
     conn: Connection
 
 
+@final
+class TraceEventData(FrozenModel):
+    name: str
+    start_us: int
+    duration_us: int
+    rank: int
+    category: str
+
+
+@final
+class TracesCollected(BaseEvent):
+    task_id: TaskId
+    rank: int
+    traces: list[TraceEventData]
+
+
+@final
+class TracesMerged(BaseEvent):
+    task_id: TaskId
+    traces: list[TraceEventData]
+
+
 Event = (
     TestEvent
     | TaskCreated
@@ -127,6 +150,8 @@ Event = (
     | InputChunkReceived
     | TopologyEdgeCreated
     | TopologyEdgeDeleted
+    | TracesCollected
+    | TracesMerged
 )
 
 
diff --git a/src/exo/worker/engines/image/pipeline/runner.py b/src/exo/worker/engines/image/pipeline/runner.py
index 698887b0..e1f65efd 100644
--- a/src/exo/worker/engines/image/pipeline/runner.py
+++ b/src/exo/worker/engines/image/pipeline/runner.py
@@ -6,6 +6,11 @@ from mflux.models.common.config.config import Config
 from mflux.utils.exceptions import StopImageGenerationException
 from tqdm import tqdm
 
+from exo.shared.constants import EXO_TRACING_ENABLED
+from exo.shared.tracing import (
+    clear_trace_buffer,
+    trace,
+)
 from exo.shared.types.worker.shards import PipelineShardMetadata
 from exo.worker.engines.image.config import ImageModelConfig
 from exo.worker.engines.image.models.base import (
@@ -324,6 +329,7 @@ class DiffusionRunner:
             capture_steps = set()
 
         self._reset_all_caches()
+        clear_trace_buffer()
 
         time_steps = tqdm(range(runtime_config.num_inference_steps))
 
@@ -465,20 +471,22 @@ class DiffusionRunner:
         if self.group is None:
             return self._single_node_step(t, config, latents, prompt_data)
         elif t < config.init_time_step + num_sync_steps:
-            return self._sync_pipeline_step(
-                t,
-                config,
-                latents,
-                prompt_data,
-            )
+            with trace(name=f"sync {t}", rank=self.rank, category="sync"):
+                return self._sync_pipeline_step(
+                    t,
+                    config,
+                    latents,
+                    prompt_data,
+                )
         else:
-            return self._async_pipeline_step(
-                t,
-                config,
-                latents,
-                prompt_data,
-                is_first_async_step=t == config.init_time_step + num_sync_steps,
-            )
+            with trace(name=f"async {t}", rank=self.rank, category="async"):
+                return self._async_pipeline_step(
+                    t,
+                    config,
+                    latents,
+                    prompt_data,
+                    is_first_async_step=t == config.init_time_step + num_sync_steps,
+                )
 
     def _single_node_step(
         self,
@@ -586,30 +594,41 @@ class DiffusionRunner:
 
         if self.has_joint_blocks:
             if not self.is_first_stage:
-                hidden_states = mx.distributed.recv(
-                    (batch_size, num_img_tokens, hidden_dim),
-                    dtype,
-                    self.prev_rank,
-                    group=self.group,
-                )
-                encoder_hidden_states = mx.distributed.recv(
-                    (batch_size, text_seq_len, hidden_dim),
-                    dtype,
-                    self.prev_rank,
-                    group=self.group,
-                )
-                mx.eval(hidden_states, encoder_hidden_states)
+                with trace(
+                    name=f"recv {self.prev_rank}", rank=self.rank, category="comms"
+                ):
+                    hidden_states = mx.distributed.recv(
+                        (batch_size, num_img_tokens, hidden_dim),
+                        dtype,
+                        self.prev_rank,
+                        group=self.group,
+                    )
+                    encoder_hidden_states = mx.distributed.recv(
+                        (batch_size, text_seq_len, hidden_dim),
+                        dtype,
+                        self.prev_rank,
+                        group=self.group,
+                    )
+                    mx.eval(hidden_states, encoder_hidden_states)
 
             assert self.joint_block_wrappers is not None
             assert encoder_hidden_states is not None
-            for wrapper in self.joint_block_wrappers:
-                wrapper.set_patch(BlockWrapperMode.CACHING)
-                encoder_hidden_states, hidden_states = wrapper(
-                    hidden_states=hidden_states,
-                    encoder_hidden_states=encoder_hidden_states,
-                    text_embeddings=text_embeddings,
-                    rotary_embeddings=image_rotary_embeddings,
-                )
+            with trace(
+                name="joint_blocks",
+                rank=self.rank,
+                category="compute",
+            ):
+                for wrapper in self.joint_block_wrappers:
+                    wrapper.set_patch(BlockWrapperMode.CACHING)
+                    encoder_hidden_states, hidden_states = wrapper(
+                        hidden_states=hidden_states,
+                        encoder_hidden_states=encoder_hidden_states,
+                        text_embeddings=text_embeddings,
+                        rotary_embeddings=image_rotary_embeddings,
+                    )
+
+                if EXO_TRACING_ENABLED:
+                    mx.eval(encoder_hidden_states, hidden_states)
 
         if self.owns_concat_stage:
             assert encoder_hidden_states is not None
@@ -620,45 +639,63 @@ class DiffusionRunner:
             if self.has_single_blocks or self.is_last_stage:
                 hidden_states = concatenated
             else:
-                concatenated = mx.distributed.send(
-                    concatenated, self.next_rank, group=self.group
-                )
-                mx.async_eval(concatenated)
+                with trace(
+                    name=f"send {self.next_rank}", rank=self.rank, category="comms"
+                ):
+                    concatenated = mx.distributed.send(
+                        concatenated, self.next_rank, group=self.group
+                    )
+                    mx.async_eval(concatenated)
 
         elif self.has_joint_blocks and not self.is_last_stage:
             assert encoder_hidden_states is not None
-            hidden_states = mx.distributed.send(
-                hidden_states, self.next_rank, group=self.group
-            )
-            encoder_hidden_states = mx.distributed.send(
-                encoder_hidden_states, self.next_rank, group=self.group
-            )
-            mx.async_eval(hidden_states, encoder_hidden_states)
+            with trace(name=f"send {self.next_rank}", rank=self.rank, category="comms"):
+                hidden_states = mx.distributed.send(
+                    hidden_states, self.next_rank, group=self.group
+                )
+                encoder_hidden_states = mx.distributed.send(
+                    encoder_hidden_states, self.next_rank, group=self.group
+                )
+                mx.async_eval(hidden_states, encoder_hidden_states)
 
         if self.has_single_blocks:
             if not self.owns_concat_stage and not self.is_first_stage:
-                hidden_states = mx.distributed.recv(
-                    (batch_size, text_seq_len + num_img_tokens, hidden_dim),
-                    dtype,
-                    self.prev_rank,
-                    group=self.group,
-                )
-                mx.eval(hidden_states)
+                with trace(
+                    name=f"recv {self.prev_rank}", rank=self.rank, category="comms"
+                ):
+                    hidden_states = mx.distributed.recv(
+                        (batch_size, text_seq_len + num_img_tokens, hidden_dim),
+                        dtype,
+                        self.prev_rank,
+                        group=self.group,
+                    )
+                    mx.eval(hidden_states)
 
             assert self.single_block_wrappers is not None
-            for wrapper in self.single_block_wrappers:
-                wrapper.set_patch(BlockWrapperMode.CACHING)
-                hidden_states = wrapper(
-                    hidden_states=hidden_states,
-                    text_embeddings=text_embeddings,
-                    rotary_embeddings=image_rotary_embeddings,
-                )
+            with trace(
+                name="single blocks",
+                rank=self.rank,
+                category="compute",
+            ):
+                for wrapper in self.single_block_wrappers:
+                    wrapper.set_patch(BlockWrapperMode.CACHING)
+                    hidden_states = wrapper(
+                        hidden_states=hidden_states,
+                        text_embeddings=text_embeddings,
+                        rotary_embeddings=image_rotary_embeddings,
+                    )
+
+                if EXO_TRACING_ENABLED:
+                    mx.eval(hidden_states)
 
             if not self.is_last_stage:
-                hidden_states = mx.distributed.send(
-                    hidden_states, self.next_rank, group=self.group
-                )
-                mx.async_eval(hidden_states)
+                with trace(
+                    name=f"send {self.next_rank}", rank=self.rank, category="comms"
+                ):
+                    hidden_states = mx.distributed.send(
+                        hidden_states, self.next_rank, group=self.group
+                    )
+                    mx.async_eval(hidden_states)
 
         hidden_states = hidden_states[:, text_seq_len:, ...]
 
@@ -742,14 +779,20 @@ class DiffusionRunner:
             )
 
             if not self.is_first_stage:
-                hidden_states = mx.distributed.send(hidden_states, 0, group=self.group)
-                mx.async_eval(hidden_states)
+                with trace(name="send 0", rank=self.rank, category="comms"):
+                    hidden_states = mx.distributed.send(
+                        hidden_states, 0, group=self.group
+                    )
+                    mx.async_eval(hidden_states)
 
         elif self.is_first_stage:
-            hidden_states = mx.distributed.recv_like(
-                prev_latents, src=self.world_size - 1, group=self.group
-            )
-            mx.eval(hidden_states)
+            with trace(
+                name=f"recv {self.world_size - 1}", rank=self.rank, category="comms"
+            ):
+                hidden_states = mx.distributed.recv_like(
+                    prev_latents, src=self.world_size - 1, group=self.group
+                )
+                mx.eval(hidden_states)
 
         else:
             hidden_states = prev_latents
@@ -809,10 +852,13 @@ class DiffusionRunner:
                 and not self.is_last_stage
                 and not is_first_async_step
             ):
-                patch = mx.distributed.recv_like(
-                    patch, src=self.prev_rank, group=self.group
-                )
-                mx.eval(patch)
+                with trace(
+                    name=f"recv {self.prev_rank}", rank=self.rank, category="comms"
+                ):
+                    patch = mx.distributed.recv_like(
+                        patch, src=self.prev_rank, group=self.group
+                    )
+                    mx.eval(patch)
 
             step_patch = mx.concatenate([patch, patch], axis=0) if needs_cfg else patch
 
@@ -843,10 +889,13 @@ class DiffusionRunner:
                 )
 
                 if not self.is_first_stage and t != config.num_inference_steps - 1:
-                    patch_latents[patch_idx] = mx.distributed.send(
-                        patch_latents[patch_idx], self.next_rank, group=self.group
-                    )
-                    mx.async_eval(patch_latents[patch_idx])
+                    with trace(
+                        name=f"send {self.next_rank}", rank=self.rank, category="comms"
+                    ):
+                        patch_latents[patch_idx] = mx.distributed.send(
+                            patch_latents[patch_idx], self.next_rank, group=self.group
+                        )
+                        mx.async_eval(patch_latents[patch_idx])
 
         return mx.concatenate(patch_latents, axis=1)
 
@@ -885,22 +934,28 @@ class DiffusionRunner:
         if self.has_joint_blocks:
             if not self.is_first_stage:
                 patch_len = patch.shape[1]
-                patch = mx.distributed.recv(
-                    (batch_size, patch_len, hidden_dim),
-                    patch.dtype,
-                    self.prev_rank,
-                    group=self.group,
-                )
-                mx.eval(patch)
-
-                if patch_idx == 0:
-                    encoder_hidden_states = mx.distributed.recv(
-                        (batch_size, text_seq_len, hidden_dim),
+                with trace(
+                    name=f"recv {self.prev_rank}", rank=self.rank, category="comms"
+                ):
+                    patch = mx.distributed.recv(
+                        (batch_size, patch_len, hidden_dim),
                         patch.dtype,
                         self.prev_rank,
                         group=self.group,
                     )
-                    mx.eval(encoder_hidden_states)
+                    mx.eval(patch)
+
+                if patch_idx == 0:
+                    with trace(
+                        name=f"recv {self.prev_rank}", rank=self.rank, category="comms"
+                    ):
+                        encoder_hidden_states = mx.distributed.recv(
+                            (batch_size, text_seq_len, hidden_dim),
+                            patch.dtype,
+                            self.prev_rank,
+                            group=self.group,
+                        )
+                        mx.eval(encoder_hidden_states)
 
             if self.is_first_stage:
                 patch, encoder_hidden_states = self.adapter.compute_embeddings(
@@ -909,14 +964,22 @@ class DiffusionRunner:
 
             assert self.joint_block_wrappers is not None
             assert encoder_hidden_states is not None
-            for wrapper in self.joint_block_wrappers:
-                wrapper.set_patch(BlockWrapperMode.PATCHED, start_token, end_token)
-                encoder_hidden_states, patch = wrapper(
-                    hidden_states=patch,
-                    encoder_hidden_states=encoder_hidden_states,
-                    text_embeddings=text_embeddings,
-                    rotary_embeddings=image_rotary_embeddings,
-                )
+            with trace(
+                name=f"joint patch {patch_idx}",
+                rank=self.rank,
+                category="compute",
+            ):
+                for wrapper in self.joint_block_wrappers:
+                    wrapper.set_patch(BlockWrapperMode.PATCHED, start_token, end_token)
+                    encoder_hidden_states, patch = wrapper(
+                        hidden_states=patch,
+                        encoder_hidden_states=encoder_hidden_states,
+                        text_embeddings=text_embeddings,
+                        rotary_embeddings=image_rotary_embeddings,
+                    )
+
+                if EXO_TRACING_ENABLED:
+                    mx.eval(encoder_hidden_states, patch)
 
         if self.owns_concat_stage:
             assert encoder_hidden_states is not None
@@ -925,49 +988,70 @@ class DiffusionRunner:
             if self.has_single_blocks or self.is_last_stage:
                 patch = patch_concat
             else:
-                patch_concat = mx.distributed.send(
-                    patch_concat, self.next_rank, group=self.group
-                )
-                mx.async_eval(patch_concat)
+                with trace(
+                    name=f"send {self.next_rank}", rank=self.rank, category="comms"
+                ):
+                    patch_concat = mx.distributed.send(
+                        patch_concat, self.next_rank, group=self.group
+                    )
+                    mx.async_eval(patch_concat)
 
         elif self.has_joint_blocks and not self.is_last_stage:
-            patch = mx.distributed.send(patch, self.next_rank, group=self.group)
-            mx.async_eval(patch)
+            with trace(name=f"send {self.next_rank}", rank=self.rank, category="comms"):
+                patch = mx.distributed.send(patch, self.next_rank, group=self.group)
+                mx.async_eval(patch)
 
             if patch_idx == 0:
                 assert encoder_hidden_states is not None
-                encoder_hidden_states = mx.distributed.send(
-                    encoder_hidden_states, self.next_rank, group=self.group
-                )
-                mx.async_eval(encoder_hidden_states)
+                with trace(
+                    name=f"send {self.next_rank}", rank=self.rank, category="comms"
+                ):
+                    encoder_hidden_states = mx.distributed.send(
+                        encoder_hidden_states, self.next_rank, group=self.group
+                    )
+                    mx.async_eval(encoder_hidden_states)
 
         if self.has_single_blocks:
             if not self.owns_concat_stage and not self.is_first_stage:
                 patch_len = patch.shape[1]
-                patch = mx.distributed.recv(
-                    (batch_size, text_seq_len + patch_len, hidden_dim),
-                    patch.dtype,
-                    self.prev_rank,
-                    group=self.group,
-                )
-                mx.eval(patch)
+                with trace(
+                    name=f"recv {self.prev_rank}", rank=self.rank, category="comms"
+                ):
+                    patch = mx.distributed.recv(
+                        (batch_size, text_seq_len + patch_len, hidden_dim),
+                        patch.dtype,
+                        self.prev_rank,
+                        group=self.group,
+                    )
+                    mx.eval(patch)
 
             assert self.single_block_wrappers is not None
-            for wrapper in self.single_block_wrappers:
-                wrapper.set_patch(BlockWrapperMode.PATCHED, start_token, end_token)
-                patch = wrapper(
-                    hidden_states=patch,
-                    text_embeddings=text_embeddings,
-                    rotary_embeddings=image_rotary_embeddings,
-                )
+            with trace(
+                name=f"single patch {patch_idx}",
+                rank=self.rank,
+                category="compute",
+            ):
+                for wrapper in self.single_block_wrappers:
+                    wrapper.set_patch(BlockWrapperMode.PATCHED, start_token, end_token)
+                    patch = wrapper(
+                        hidden_states=patch,
+                        text_embeddings=text_embeddings,
+                        rotary_embeddings=image_rotary_embeddings,
+                    )
+
+                if EXO_TRACING_ENABLED:
+                    mx.eval(patch)
 
             if not self.is_last_stage:
-                patch = mx.distributed.send(patch, self.next_rank, group=self.group)
-                mx.async_eval(patch)
+                with trace(
+                    name=f"send {self.next_rank}", rank=self.rank, category="comms"
+                ):
+                    patch = mx.distributed.send(patch, self.next_rank, group=self.group)
+                    mx.async_eval(patch)
 
         noise: mx.array | None = None
         if self.is_last_stage:
-            patch_img_only = patch[:, text_seq_len:, :]
-            noise = self.adapter.final_projection(patch_img_only, text_embeddings)
+            patch = patch[:, text_seq_len:, :]
+            noise = self.adapter.final_projection(patch, text_embeddings)
 
         return noise, encoder_hidden_states
diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py
index 0bc2e846..0d012fb6 100644
--- a/src/exo/worker/runner/runner.py
+++ b/src/exo/worker/runner/runner.py
@@ -16,8 +16,9 @@ from openai_harmony import (  # pyright: ignore[reportMissingTypeStubs]
 )
 from pydantic import ValidationError
 
-from exo.shared.constants import EXO_MAX_CHUNK_SIZE
+from exo.shared.constants import EXO_MAX_CHUNK_SIZE, EXO_TRACING_ENABLED
 from exo.shared.models.model_cards import ModelId, ModelTask
+from exo.shared.tracing import clear_trace_buffer, get_trace_buffer
 from exo.shared.types.api import ImageGenerationStats
 from exo.shared.types.chunks import ErrorChunk, ImageChunk, TokenChunk, ToolCallChunk
 from exo.shared.types.common import CommandId
@@ -27,6 +28,8 @@ from exo.shared.types.events import (
     RunnerStatusUpdated,
     TaskAcknowledged,
     TaskStatusUpdated,
+    TraceEventData,
+    TracesCollected,
 )
 from exo.shared.types.tasks import (
     ConnectToGroup,
@@ -408,6 +411,10 @@ def main(
                                 )
                             )
                         raise
+                    finally:
+                        _send_traces_if_enabled(
+                            event_sender, task.task_id, shard_metadata.device_rank
+                        )
 
                     current_status = RunnerReady()
                     logger.info("runner ready")
@@ -466,6 +473,10 @@ def main(
                                 )
                             )
                         raise
+                    finally:
+                        _send_traces_if_enabled(
+                            event_sender, task.task_id, shard_metadata.device_rank
+                        )
 
                     current_status = RunnerReady()
                     logger.info("runner ready")
@@ -640,6 +651,36 @@ def _send_image_chunk(
         )
 
 
+def _send_traces_if_enabled(
+    event_sender: MpSender[Event],
+    task_id: TaskId,
+    rank: int,
+) -> None:
+    if not EXO_TRACING_ENABLED:
+        return
+
+    traces = get_trace_buffer()
+    if traces:
+        trace_data = [
+            TraceEventData(
+                name=t.name,
+                start_us=t.start_us,
+                duration_us=t.duration_us,
+                rank=t.rank,
+                category=t.category,
+            )
+            for t in traces
+        ]
+        event_sender.send(
+            TracesCollected(
+                task_id=task_id,
+                rank=rank,
+                traces=trace_data,
+            )
+        )
+    clear_trace_buffer()
+
+
 def _process_image_response(
     response: ImageGenerationResponse | PartialImageResponse,
     command_id: CommandId,

← 3e27ead7 remove mdns discovered peers from appearing in state (#1312  ·  back to Exo  ·  create config home when checking for config file (#1353) 54b19879 →