diff --git a/miles/dashboard/collector.py b/miles/dashboard/collector.py index 245ba30e..989accb8 100644 --- a/miles/dashboard/collector.py +++ b/miles/dashboard/collector.py @@ -7,8 +7,8 @@ from dataclasses import dataclass, field from typing import Any, ClassVar -from miles.dashboard.events import PhaseEvent, TrajectoryEvent -from miles.dashboard.store import DashboardStore, Record, Stream +from miles.dashboard.events import PhaseEvent, RequestEvent, TrajectoryEvent +from miles.dashboard.store import DashboardStore, Record, stream_for logger = logging.getLogger(__name__) @@ -59,8 +59,12 @@ def push_trajectories(self, batch: list[TrajectoryEvent]) -> None: for event in batch: self._append(event) + def push_requests(self, batch: list[RequestEvent]) -> None: + for event in batch: + self._append(event) + def _append(self, record: Record) -> None: - stream = Stream.PHASES if isinstance(record, PhaseEvent) else Stream.TRAJECTORIES + stream = stream_for(record) with self._lock: if self._store.buffered_count(stream) >= self.MAX_BUFFERED_PER_STREAM: self._dropped_since_flush += self._store.drop_oldest_buffered(stream) diff --git a/miles/dashboard/events.py b/miles/dashboard/events.py index f5559f61..49ef202f 100644 --- a/miles/dashboard/events.py +++ b/miles/dashboard/events.py @@ -53,6 +53,28 @@ class TrajectoryEventKind(StrEnum): } +@dataclass +class RequestEvent: + """One rollout request's wall-clock marks, keyed by miles.utils.request_timing. + + The marks, not the leg durations derived from them: the reader places every + leg at the time it happened, and durations are a subtraction away. + ``engine_stages`` is the engine's own breakdown of the one leg it owns. + """ + + rollout_id: int + request_id: str + group_index: int + sample_indices: list[int] + marks: dict[str, float] + engine_stages: dict[str, float] + worker: str + resp_bytes: int + + def to_dict(self) -> dict: + return asdict(self) + + @dataclass class TrajectoryEvent: ts: float diff --git a/miles/dashboard/hooks.py b/miles/dashboard/hooks.py index e4566d65..62c1db84 100644 --- a/miles/dashboard/hooks.py +++ b/miles/dashboard/hooks.py @@ -10,7 +10,18 @@ from contextlib import contextmanager from pathlib import Path -from miles.dashboard.events import STAGE_KINDS, PhaseEvent, TrajectoryEvent +from miles.dashboard.events import STAGE_KINDS, PhaseEvent, RequestEvent, TrajectoryEvent +from miles.utils.request_timing import ( + ROUTER_TIMING_HEADER, + ROUTER_WIRE_KEYS, + ROUTER_WORKER_HEADER, + SGLD_STAGES_HEADER, + SGLD_TIMING_HEADER, + SGLD_WIRE_KEYS, + Marks, + parse_header, + parse_stages, +) logger = logging.getLogger(__name__) @@ -119,8 +130,82 @@ def flush(self) -> None: logger.warning("dashboard trajectory sink flush failed; dropping events", exc_info=True) +class RequestSink: + def __init__(self, handle) -> None: + self.handle = handle + self._buffer: list[RequestEvent] = [] + self._lock = threading.Lock() + self._last_flush = time.monotonic() + + def record(self, tracer: RequestTracer, samples) -> None: + try: + sample = samples[0] + event = RequestEvent( + rollout_id=tracer.rollout_id, + request_id=sample.request_id or "", + group_index=sample.group_index if sample.group_index is not None else -1, + sample_indices=[s.index if s.index is not None else -1 for s in samples], + marks={name: round(ts, 6) for name, ts in tracer.marks.marks.items()}, + engine_stages=tracer.engine_stages, + worker=tracer.worker, + resp_bytes=tracer.resp_bytes, + ) + with self._lock: + self._buffer.append(event) + batch = self._take_batch_if_due() + if batch: + self.handle.push_requests.remote(batch) + except Exception: # noqa: BLE001 + logger.warning("dashboard request sink failed; dropping events", exc_info=True) + + def _take_batch_if_due(self) -> list[RequestEvent] | None: + if len(self._buffer) < BATCH_MAX_EVENTS and time.monotonic() - self._last_flush < BATCH_MAX_SECONDS: + return None + batch, self._buffer = self._buffer, [] + self._last_flush = time.monotonic() + return batch + + def flush(self) -> None: + try: + with self._lock: + batch, self._buffer = self._buffer, [] + if batch: + _ray_get(self.handle.push_requests.remote(batch)) + except Exception: # noqa: BLE001 + logger.warning("dashboard request sink flush failed; dropping events", exc_info=True) + + +class RequestTracer: + """Collects one rollout request's marks; ``done`` is a no-op with the + dashboard off, so the rollout path can carry a tracer unconditionally.""" + + def __init__(self, rollout_id: int) -> None: + self.rollout_id = rollout_id + self.marks = Marks() + self.engine_stages: dict[str, float] = {} + self.worker = "" + self.resp_bytes = 0 + + def mark(self, name: str) -> None: + self.marks.mark(name) + + def absorb_response(self, result) -> None: + """Take the sender's own marks plus the ones the reply carries back.""" + self.marks.absorb(result.marks) + headers = {key.lower(): value for key, value in result.headers.items()} + self.marks.absorb(parse_header(headers.get(SGLD_TIMING_HEADER), SGLD_WIRE_KEYS)) + self.marks.absorb(parse_header(headers.get(ROUTER_TIMING_HEADER), ROUTER_WIRE_KEYS)) + self.engine_stages = parse_stages(headers.get(SGLD_STAGES_HEADER)) + self.worker = headers.get(ROUTER_WORKER_HEADER, "") + + def done(self, samples) -> None: + if _request_sink is not None: + _request_sink.record(self, samples) + + _phase_sink: PhaseSink | None = None _trajectory_sink: TrajectorySink | None = None +_request_sink: RequestSink | None = None _rollout_id = -1 _GPU_SAMPLER: GpuUtilSampler | None = None @@ -152,6 +237,7 @@ def register_rollout_manager(args) -> None: return attach_phase_sink(handle, "rollout") attach_trajectory_sink(handle) + attach_request_sink(handle) def set_rollout_id(rollout_id: int) -> None: @@ -159,6 +245,10 @@ def set_rollout_id(rollout_id: int) -> None: _rollout_id = rollout_id +def current_rollout_id() -> int: + return _rollout_id + + def record_trajectory(sample) -> None: if _trajectory_sink is not None: _trajectory_sink.record(sample, _rollout_id) @@ -201,8 +291,14 @@ def attach_trajectory_sink(handle) -> None: _trajectory_sink = TrajectorySink(handle) +def attach_request_sink(handle) -> None: + global _request_sink + if _request_sink is None: + _request_sink = RequestSink(handle) + + def detach_and_flush() -> None: - global _phase_sink, _trajectory_sink, _GPU_SAMPLER + global _phase_sink, _trajectory_sink, _request_sink, _GPU_SAMPLER from miles.utils.timer import Timer if _phase_sink is not None: @@ -212,6 +308,9 @@ def detach_and_flush() -> None: if _trajectory_sink is not None: _trajectory_sink.flush() _trajectory_sink = None + if _request_sink is not None: + _request_sink.flush() + _request_sink = None if _GPU_SAMPLER is not None: _GPU_SAMPLER.stop() _GPU_SAMPLER = None diff --git a/miles/dashboard/store.py b/miles/dashboard/store.py index c08dc869..e8462cf9 100644 --- a/miles/dashboard/store.py +++ b/miles/dashboard/store.py @@ -8,7 +8,7 @@ from enum import StrEnum from pathlib import Path -from miles.dashboard.events import PhaseEvent, TrajectoryEvent +from miles.dashboard.events import PhaseEvent, RequestEvent, TrajectoryEvent RUN_DIR_PREFIX = "run_" @@ -17,20 +17,25 @@ class Stream(StrEnum): PHASES = "phases" TRAJECTORIES = "trajectories" + REQUESTS = "requests" -Record = PhaseEvent | TrajectoryEvent +Record = PhaseEvent | TrajectoryEvent | RequestEvent -def _stream(record: Record) -> Stream: +def stream_for(record: Record) -> Stream: if isinstance(record, PhaseEvent): return Stream.PHASES + if isinstance(record, RequestEvent): + return Stream.REQUESTS return Stream.TRAJECTORIES def _timestamp(record: Record) -> float: if isinstance(record, PhaseEvent): return record.t1 + if isinstance(record, RequestEvent): + return max(record.marks.values(), default=0.0) return record.ts @@ -63,7 +68,7 @@ def write_meta(self, *, run_name: str, start_ts: float, args: dict) -> None: ) def append(self, record: Record) -> None: - self._buffers[_stream(record)].append(record) + self._buffers[stream_for(record)].append(record) def buffered_count(self, stream: Stream) -> int: return len(self._buffers[stream]) diff --git a/miles/dashboard/viewer.py b/miles/dashboard/viewer.py index ca9384b1..6908023f 100644 --- a/miles/dashboard/viewer.py +++ b/miles/dashboard/viewer.py @@ -13,6 +13,23 @@ from miles.dashboard.events import SPAN_KINDS from miles.dashboard.store import resolve_run_dir +from miles.utils.request_timing import CROSS, LEG_NAMES, LEGS, derive, leg_sources + +# one hue per process, so a waterfall reads first as "where was the request"; +# the cross-process hops stay near-neutral, distinguished by lightness alone +_SOURCE_HUE = {"client": (212, 60), "router": (22, 62), "sgld": (158, 55), CROSS: (35, 10)} + + +def _leg_colors() -> dict[str, str]: + sources = leg_sources() + seen: dict[str, int] = {} + colors = {} + for name in LEG_NAMES: + hue, saturation = _SOURCE_HUE[sources[name]] + index = seen.get(sources[name], 0) + seen[sources[name]] = index + 1 + colors[name] = f"hsl({hue},{saturation - index * 3}%,{38 + min(index, 6) * 6}%)" + return colors def _read_jsonl(paths): @@ -54,17 +71,26 @@ def _fold_trajectory(events): return segments -def load_streams(workspace: str): +def load_streams(workspace: str, max_rollouts: int = 0): phases = _read_jsonl(sorted(glob.glob(os.path.join(workspace, "phases", "*.jsonl")))) gpu = _read_jsonl(sorted(glob.glob(os.path.join(workspace, "gpu_util", "*.jsonl")))) traj = _read_jsonl(sorted(glob.glob(os.path.join(workspace, "trajectories", "*.jsonl")))) - return phases, gpu, _fold_trajectory(traj) + reqs = _read_jsonl(sorted(glob.glob(os.path.join(workspace, "requests", "*.jsonl")))) + if max_rollouts: + # the page embeds every record, so a long run needs a window + keep = sorted({r.get("rollout_id", -1) for r in reqs})[-max_rollouts:] + reqs = [r for r in reqs if r.get("rollout_id", -1) in keep] + return phases, gpu, _fold_trajectory(traj), reqs -def _compute_data(phases, gpu, life=None) -> dict: +def _compute_data(phases, gpu, life=None, reqs=None) -> dict: life = life or [] + timings = [(r, derive(r.get("marks") or {})) for r in (reqs or [])] + timings = [(r, t) for r, t in timings if t.t_end > t.t_start] # normalize all timestamps to the earliest event across streams - ts_candidates = [p["t0"] for p in phases] + [g["ts"] for g in gpu] + [x["t0"] for x in life] + ts_candidates = ( + [p["t0"] for p in phases] + [g["ts"] for g in gpu] + [x["t0"] for x in life] + [t.t_start for _, t in timings] + ) t0 = min(ts_candidates) if ts_candidates else 0.0 ph = [ { @@ -98,12 +124,36 @@ def _compute_data(phases, gpu, life=None) -> dict: for x in life if x.get("t1", 0) >= x.get("t0", 0) ] - span = max([r["e"] for r in ph] + [r["t"] for r in gp] + [r["e"] for r in lf] + [1.0]) - return {"phases": ph, "gpu": gp, "life": lf, "span": span} + rq = [ + { + "k": (r.get("request_id") or "?")[:12], + "r": r.get("rollout_id", -1), + "s": round(timing.t_start - t0, 3), + "e": round(timing.t_end - t0, 3), + "x": round(timing.cross_total, 4), + "w": r.get("worker", ""), + "n": len(r.get("sample_indices") or []), + "m": {name: round(ts - t0, 4) for name, ts in (r.get("marks") or {}).items()}, + "g": {name: round(sec, 4) for name, sec in (r.get("engine_stages") or {}).items()}, + } + for r, timing in timings + ] + rq.sort(key=lambda r: r["s"]) + span = max([r["e"] for r in ph] + [r["t"] for r in gp] + [r["e"] for r in lf] + [r["e"] for r in rq] + [1.0]) + return { + "phases": ph, + "gpu": gp, + "life": lf, + "reqs": rq, + "legs": [list(leg) for leg in LEGS], + "legSource": leg_sources(), + "legColors": _leg_colors(), + "span": span, + } -def build_html(phases, gpu, life=None, title="miles-D rollout dashboard") -> str: - data = json.dumps(_compute_data(phases, gpu, life), separators=(",", ":")) +def build_html(phases, gpu, life=None, reqs=None, title="miles-D rollout dashboard") -> str: + data = json.dumps(_compute_data(phases, gpu, life, reqs), separators=(",", ":")) return _TEMPLATE.replace("__TITLE__", title).replace("__SERVE__", "false").replace("__DATA__", data) @@ -141,6 +191,13 @@ def build_html(phases, gpu, life=None, title="miles-D rollout dashboard") -> str .grid{stroke:var(--border);stroke-width:1;} .rowlab{fill:var(--text);font:11px ui-monospace,monospace;} .roundline{stroke:var(--accent);stroke-width:1;stroke-dasharray:3 3;opacity:.5;} +table.pct{border-collapse:collapse;width:100%;font-size:12px;font-family:ui-monospace,monospace;} +table.pct th,table.pct td{text-align:right;padding:3px 8px;border-bottom:1px solid var(--border);} +table.pct th{color:var(--muted);font-weight:600;} +table.pct th.leg,table.pct td.leg{text-align:left;} +.sw{width:10px;height:10px;border-radius:2px;display:inline-block;margin-right:6px;vertical-align:-1px;} +select{background:var(--bg);color:var(--text);border:1px solid var(--border);border-radius:4px;padding:3px 6px;font:inherit;font-size:12px;} +.warn{color:var(--accent);} #tooltip{position:fixed;pointer-events:none;background:var(--topbar-bg);color:var(--topbar-text);border:1px solid var(--topbar-border);border-radius:4px;padding:6px 9px;font-size:12px;z-index:100;opacity:0;font-family:ui-monospace,monospace;white-space:pre;line-height:1.5;}
@@ -156,6 +213,21 @@ def build_html(phases, gpu, life=None, title="miles-D rollout dashboard") -> str

GPU utilization %

Phase timeline (cross-round)

+
+
+

Request waterfall

+
+ + +
+
+ +
+

Where a request's time goes (mean per request, per rollout)

+
+

Leg latency

+

Engine stages inside sgld_forward

+

Per-sample lifecycle

@@ -180,11 +252,16 @@ def build_html(phases, gpu, life=None, title="miles-D rollout dashboard") -> str : phaseNames.includes("actor_train") ? "actor_train" : (phaseNames[0]||""); rounds = D.phases.filter(p=>p.name===boundaryName).map(p=>p.s).sort((a,b)=>a-b); if(!userZoomed){ x0=0; x1=D.span; } else if(x1>D.span){ x1=D.span; } + if(!D.reqs) D.reqs=[]; + rebuildReqs(); const lg=document.getElementById("lifelegend"); - if(D.life.length){ lg.innerHTML=lifeStages.map(s=>`${s}`).join(""); } - else{ document.getElementById("samples").closest(".panel").style.display="none"; } + // the waterfall is the same data per request, so the coarse panel is a fallback + const showLife = D.life.length && !D.reqs.length; + if(showLife){ lg.innerHTML=lifeStages.map(s=>`${s}`).join(""); } + document.getElementById("samples").closest(".panel").style.display = showLife ? "" : "none"; document.getElementById("runinfo").textContent = - `span ${D.span.toFixed(1)}s · ${D.phases.length} phases · ${gpuIds.length} GPUs · ${D.gpu.length} gpu samples · ${sampKeys.length} reqs`; + `span ${D.span.toFixed(1)}s · ${D.phases.length} phases · ${gpuIds.length} GPUs · ${D.gpu.length} gpu samples` + + ` · ${D.reqs.length ? D.reqs.length + " requests" : sampKeys.length + " reqs"}`; const rb=document.getElementById("rounds"); rb.innerHTML=""; if(rounds.length){ rb.insertAdjacentHTML("beforeend",`rounds (${boundaryName}): `); @@ -200,7 +277,8 @@ def build_html(phases, gpu, life=None, title="miles-D rollout dashboard") -> str function niceStep(r){const p=Math.pow(10,Math.floor(Math.log10(r)));const f=r/p;return (f<1.5?1:f<3?2:f<7?5:10)*p;} function draw(){ - drawGpu(); drawGantt(); if(D.life.length)drawSamples(); + drawGpu(); drawGantt(); + if(D.reqs.length) drawWaterfall(); else if(D.life.length) drawSamples(); } function drawSamples(){ const svg=document.getElementById("samples"),w=W; @@ -230,6 +308,162 @@ def build_html(phases, gpu, life=None, title="miles-D rollout dashboard") -> str el.addEventListener("mouseleave",()=>tip.style.opacity=0); }); } +const WF_ROW=4, WF_GAP=1, WF_MAX=2000; +let wfShown=[]; + +function selectedReqs(){ + const v=document.getElementById("wfroll").value; + return v==="all" ? D.reqs : D.reqs.filter(r=>String(r.r)===v); +} +// every leg is the span between two consecutive marks, so it draws at the time +// it happened; a leg whose marks come from two processes carries their clock +// offset, which is why an inverted one is flagged rather than hidden +function legSpans(r){ + const out=[]; + for(const [name,a,b] of D.legs){ + if(r.m[a]===undefined||r.m[b]===undefined) continue; + out.push([name,r.m[a],r.m[b]]); + } + return out; +} +function activeLegs(){ return D.legs.map(l=>l[0]).filter(n=>D.reqs.some(r=>r._d[n]!==undefined)); } +function rebuildReqs(){ + const box=document.getElementById("reqpanels"); + box.style.display = D.reqs.length ? "" : "none"; + if(!D.reqs.length) return; + D.reqs.forEach(r=>{ + r._legs=legSpans(r); + r._d={}; + r._legs.forEach(([n,a,b])=>{r._d[n]=b-a;}); + r._skew=r._legs.some(([n,a,b])=>br.r))].sort((a,b)=>a-b); + sel.innerHTML=ids.map(i=>``).join("")+``; + sel.value=[...sel.options].some(o=>o.value===prev)?prev:String(ids[ids.length-1]); + document.getElementById("wflegend").innerHTML=activeLegs() + .map(l=>`${l}`).join(""); + const skew=D.reqs.filter(r=>r._skew).length; + document.getElementById("wfwarn").innerHTML = skew + ? `${skew} request(s) have an inverted cross-process hop: those node clocks disagree, so a remote block is drawn shifted.` + : ""; +} +function reqTip(r){ + const total=r.e-r.s; + const rows=Object.keys(r._d).sort((a,b)=>r._d[b]-r._d[a]).map(n=> + ` ${n.padEnd(22)}${(r._d[n]*1000).toFixed(0).padStart(8)}ms${(100*r._d[n]/(total||1)).toFixed(0).padStart(5)}%` + +(D.legSource[n]==="cross"?" †":"")); + return [`${r.k} · R${r.r} · ${r.n} sample(s)`, r.w?`worker ${r.w}`:"", + `total ${total.toFixed(2)}s · cross-process ${r.x.toFixed(3)}s exact`, ...rows, + " † own value carries a clock offset; only their sum is exact"] + .filter(Boolean).join("\n"); +} +function drawWaterfall(){ + const svg=document.getElementById("waterfall"),w=W; + wfShown=selectedReqs().slice(0,WF_MAX); + const bandTop=18,h=bandTop+Math.max(wfShown.length,1)*(WF_ROW+WF_GAP)+30; + svg.setAttribute("viewBox",`0 0 ${w} ${h}`); + let s=axis(svg,w,h); + s+=``; + wfShown.forEach((r,i)=>{ + const y=bandTop+i*(WF_ROW+WF_GAP); + for(const [name,t0,t1] of r._legs){ + let xa=sx(t0,w),xb=sx(t1,w); + if(xbw-PADR) continue; + xa=Math.max(xa,PADL);xb=Math.min(xb,w-PADR); + s+=``; + } + }); + const more=selectedReqs().length-wfShown.length; + s+=`${wfShown.length} req`; + if(more>0) s+=`+${more} more not drawn`; + svg.innerHTML=s; + svg.querySelectorAll(".wf").forEach(el=>{ + el.addEventListener("mousemove",e=>{ + tip.textContent=reqTip(wfShown[+el.dataset.i]); + tip.style.opacity=1;tip.style.left=(e.clientX+12)+"px";tip.style.top=(e.clientY+12)+"px";}); + el.addEventListener("mouseleave",()=>tip.style.opacity=0); + }); +} +function pctile(sorted,p){ + if(!sorted.length) return 0; + return sorted[Math.min(sorted.length-1,Math.max(0,Math.ceil(p*sorted.length)-1))]; +} +function statTable(keys,valueOf,total,colorOf,label,note){ + const ms=v=>(v*1000).toFixed(v*1000<10?1:0); + const stats=keys.map(k=>{ + const vs=valueOf(k).sort((a,b)=>a-b); + if(!vs.length) return null; + const sum=vs.reduce((a,b)=>a+b,0); + return {k,n:vs.length,p50:pctile(vs,.5),p90:pctile(vs,.9),p99:pctile(vs,.99), + max:vs[vs.length-1],mean:sum/vs.length,sum,share:total?sum/total:0}; + }).filter(Boolean).sort((a,b)=>b.sum-a.sum); + if(!stats.length) return ""; + return `` + +`` + +stats.map(r=>`` + +`` + +``).join("") + +`
${label}np50p90p99maxmeanshare
${r.k}${r.n}${ms(r.p50)}${ms(r.p90)}${ms(r.p99)}${ms(r.max)}${ms(r.mean)}${(100*r.share).toFixed(1)}%
${note}
`; +} +function drawPct(){ + const rows=selectedReqs(),box=document.getElementById("pct"); + if(!rows.length){box.innerHTML="";return;} + const total=rows.reduce((a,r)=>a+(r.e-r.s),0); + const crossMean=rows.reduce((a,r)=>a+r.x,0)/rows.length; + box.innerHTML=statTable( + D.legs.map(l=>l[0]+(D.legSource[l[0]]==="cross"?" †":"")), + k=>rows.map(r=>r._d[k.replace(" †","")]).filter(v=>v!==undefined), + total, k=>D.legColors[k.replace(" †","")], "leg", + `ms per request · ${rows.length} requests · share = of summed request time` + +`
† each cross-process hop's own value carries the two clocks' offset; their sum is exact:` + +` ${(crossMean*1000).toFixed(0)}ms mean per request`); + + // The engine's own breakdown of sgld_forward, which the client otherwise + // sees as one number. Durations only: they have no marks to place them by. + const stageKeys=[...new Set(rows.flatMap(r=>Object.keys(r.g)))]; + const forward=rows.reduce((a,r)=>a+(r._d.sgld_forward||0),0); + document.getElementById("pctengine").innerHTML=statTable( + stageKeys, k=>rows.map(r=>r.g[k]).filter(v=>v!==undefined), + forward, ()=>D.legColors.sgld_forward, "engine stage", + `ms per request · share = of summed sgld_forward · reported by the engine, not placed on the axis`); +} +function drawStack(){ + const svg=document.getElementById("stack"),w=W,h=210; + const byR={}; + D.reqs.forEach(r=>{(byR[r.r]=byR[r.r]||[]).push(r);}); + const ids=Object.keys(byR).map(Number).sort((a,b)=>a-b),legs=activeLegs(); + const means=ids.map(id=>{const rs=byR[id],m={}; + legs.forEach(l=>{m[l]=rs.reduce((a,r)=>a+Math.max(r._d[l]||0,0),0)/rs.length;});return m;}); + const ymax=Math.max(...means.map(m=>legs.reduce((a,l)=>a+m[l],0)),0.001); + svg.setAttribute("viewBox",`0 0 ${w} ${h}`); + const plotH=h-30-8,slot=(w-PADL-PADR)/Math.max(ids.length,1),bw=Math.max(2,Math.min(48,slot*0.7)); + let s=``; + for(const f of [0,0.5,1]){ + const py=8+plotH-f*plotH; + s+=`` + +`${(ymax*f).toFixed(1)}s`; + } + ids.forEach((id,i)=>{ + const cx=PADL+(i+0.5)*slot; + let acc=0; + legs.forEach(l=>{ + const v=means[i][l]; + if(!v) return; + const y1=8+plotH-(acc+v)/ymax*plotH,y0=8+plotH-acc/ymax*plotH; + acc+=v; + s+=``; + }); + if(slot>24||i%Math.ceil(ids.length/24)===0) s+=`R${id}`; + }); + svg.innerHTML=s; + svg.querySelectorAll(".sb").forEach(el=>{ + el.addEventListener("mousemove",e=>{ + tip.textContent=el.dataset.t; + tip.style.opacity=1;tip.style.left=(e.clientX+12)+"px";tip.style.top=(e.clientY+12)+"px";}); + el.addEventListener("mouseleave",()=>tip.style.opacity=0); + }); +} function axis(svg,w,h){ let s=``; for(const t of ticks()){const px=sx(t,w);s+=``+ @@ -283,18 +517,20 @@ def build_html(phases, gpu, life=None, title="miles-D rollout dashboard") -> str window.addEventListener("mousemove",e=>{if(!drag)return;userZoomed=true;const w=W;const dt=(e.clientX-drag.px)/svg.clientWidth*w/(w-PADL-PADR)*(x1-x0); x0-=dt;x1-=dt;if(x0<0){x1-=x0;x0=0;}if(x1>D.span){x0-=x1-D.span;x1=D.span;}drag.px=e.clientX;draw();}); } -["gpu","gantt","samples"].forEach(id=>attachZoom(document.getElementById(id))); +["gpu","gantt","samples","waterfall"].forEach(id=>attachZoom(document.getElementById(id))); +document.getElementById("wfroll").onchange=()=>{drawWaterfall();drawPct();}; document.getElementById("reset").onclick=()=>{userZoomed=false;x0=0;x1=D.span;draw();}; async function refresh(){ if(SERVE){ try{ D=await (await fetch("data")).json(); }catch(e){ return; } } rebuild(); draw(); + if(D.reqs.length){ drawPct(); drawStack(); } } refresh(); if(SERVE) setInterval(refresh, 3000); """ -def serve(workspace: str, host: str, port: int, title: str) -> None: +def serve(workspace: str, host: str, port: int, title: str, max_rollouts: int = 0) -> None: """Live server: serves the page shell once and a fresh /data JSON each poll, so the page auto-updates a running run's data without losing the zoom view.""" import http.server @@ -311,7 +547,7 @@ def _send(self, body: bytes, ctype: str) -> None: def do_GET(self): if self.path.startswith("/data"): - body = json.dumps(_compute_data(*load_streams(workspace))).encode() + body = json.dumps(_compute_data(*load_streams(workspace, max_rollouts))).encode() self._send(body, "application/json") else: self._send(shell, "text/html; charset=utf-8") @@ -331,19 +567,28 @@ def main(): ap.add_argument( "--serve", action="store_true", help="run a live auto-updating server instead of writing a static file" ) + ap.add_argument( + "--max-rollouts", + type=int, + default=0, + help="keep only the newest N rollouts of per-request timing (0 = all)", + ) ap.add_argument("--host", default="0.0.0.0") ap.add_argument("--port", type=int, default=8000) args = ap.parse_args() workspace = str(resolve_run_dir(args.workspace)) if args.serve: - serve(workspace, args.host, args.port, args.title) + serve(workspace, args.host, args.port, args.title, args.max_rollouts) return - phases, gpu, life = load_streams(workspace) - html = build_html(phases, gpu, life, title=args.title) + phases, gpu, life, reqs = load_streams(workspace, args.max_rollouts) + html = build_html(phases, gpu, life, reqs, title=args.title) with open(args.out, "w") as f: f.write(html) - print(f"wrote {args.out}: {len(phases)} phase spans, {len(gpu)} gpu samples, {len(life)} lifecycle segs") + print( + f"wrote {args.out}: {len(phases)} phase spans, {len(gpu)} gpu samples, " + f"{len(life)} lifecycle segs, {len(reqs)} request timings" + ) if __name__ == "__main__": diff --git a/miles/rollout/sglang_diffusion_rollout.py b/miles/rollout/sglang_diffusion_rollout.py index cde28da0..306141c5 100644 --- a/miles/rollout/sglang_diffusion_rollout.py +++ b/miles/rollout/sglang_diffusion_rollout.py @@ -162,11 +162,17 @@ def submit_generate_tasks(self, samples: list[list[Sample]]) -> None: async def generate_microgroup( - args: Namespace, microgroup: list[Sample], sampling_params: dict[str, Any], *, evaluation: bool = False + args: Namespace, + microgroup: list[Sample], + sampling_params: dict[str, Any], + *, + evaluation: bool = False, + tracer: hooks.RequestTracer | None = None, ) -> list[Sample]: """Generate using traditional SGLang router with token-based workflow""" state = GenerateState(args) + tracer = tracer if tracer is not None else hooks.RequestTracer(hooks.current_rollout_id()) url = f"http://{args.sglang_router_ip}:{args.sglang_router_port}/rollout/generate" # Prepare payload for sglang-diffusion server @@ -193,10 +199,23 @@ async def generate_microgroup( st = hooks.StageTimer() with st.stage("generate"): - raw = await post(url, payload, raw=True) + # timed: the reply carries the router's and the engine's marks, the only + # way to see inside this wait + result = await post(url, payload, raw=True, timed=True) + tracer.absorb_response(result) + tracer.resp_bytes = len(result.body) with st.stage("deserialize"): - ref = state.next_parser().apply_raw.remote(microgroup, raw) - microgroup = await asyncio.to_thread(ray.get, ref) + tracer.mark("parser_submit") + # .remote() serialises the ~1GB body into plasma in the calling thread, + # so submitting on the loop blocks every other request's transfer for the + # duration of that copy. Submit and collect in the same worker thread. + parser, body = state.next_parser(), result.body + pending = microgroup + microgroup, parser_marks = await asyncio.to_thread( + lambda: ray.get(parser.apply_raw.remote(pending, body)) + ) + tracer.marks.absorb(parser_marks) + tracer.mark("parser_done") st.attach(microgroup) # Stash the SDE/training step indices on each sample so _train_core can @@ -223,9 +242,12 @@ async def generate_and_rm_microgroup( ) -> list[Sample]: state = GenerateState(args) + tracer = hooks.RequestTracer(hooks.current_rollout_id()) + tracer.mark("req_start") # generate async with state.semaphore: + tracer.mark("slot_acquired") with state.dp_rank_context() as _: if args.custom_generate_function_path is not None: custom_generate_func = load_function(args.custom_generate_function_path) @@ -234,20 +256,26 @@ async def generate_and_rm_microgroup( else: microgroup = await custom_generate_func(args, microgroup, sampling_params) else: - microgroup = await generate_microgroup(args, microgroup, sampling_params, evaluation=evaluation) + microgroup = await generate_microgroup( + args, microgroup, sampling_params, evaluation=evaluation, tracer=tracer + ) # for the rm that need the whole group, we will not do the rm here if args.group_rm: + tracer.done(microgroup) return microgroup # calculate the reward for the microgroup st = hooks.StageTimer() + tracer.mark("reward_start") with st.stage("reward"): rewards = await batched_async_rm(args, microgroup) + tracer.mark("reward_end") st.attach(microgroup) for sample, reward in zip(microgroup, rewards, strict=True): sample.reward = reward hooks.record_trajectory(sample) + tracer.done(microgroup) return microgroup @@ -278,9 +306,13 @@ async def generate_and_rm_group( # for the rm that need the whole group, we will do the rm here if args.group_rm: - rewards = await batched_async_rm(args, group) + st = hooks.StageTimer() + with st.stage("reward"): + rewards = await batched_async_rm(args, group) + st.attach(group) for sample, reward in zip(group, rewards, strict=False): sample.reward = reward + hooks.record_trajectory(sample) return group diff --git a/miles/router/router.py b/miles/router/router.py index 935ae060..f17d14bb 100644 --- a/miles/router/router.py +++ b/miles/router/router.py @@ -8,8 +8,9 @@ import uvicorn from fastapi import FastAPI, Request from fastapi.responses import JSONResponse -from starlette.responses import Response +from starlette.responses import StreamingResponse +from miles.utils.request_timing import ROUTER_TIMING_HEADER, ROUTER_WIRE_KEYS, ROUTER_WORKER_HEADER, Marks logger = logging.getLogger(__name__) @@ -126,6 +127,8 @@ async def _health_check_loop(self): async def proxy(self, request: Request, path: str): """Proxy all other requests to the SGLang router""" + stamps = Marks(ROUTER_WIRE_KEYS) + stamps.mark("router_recv") # Forward all other paths to SGLang router worker_url = self._use_url() url = f"{worker_url}/{path}" @@ -135,17 +138,34 @@ async def proxy(self, request: Request, path: str): headers = dict(request.headers) try: - response = await self.client.request(request.method, url, content=body, headers=headers) - # Pass through raw bytes — JSON re-serialization is too expensive for large tensor payloads. - content = await response.aread() - return Response( - content=content, - status_code=response.status_code, - headers=dict(response.headers), - ) - - finally: + stamps.mark("router_dispatch") + upstream = self.client.build_request(request.method, url, content=body, headers=headers) + response = await self.client.send(upstream, stream=True) + except Exception: self._finish_url(worker_url) + raise + stamps.mark("worker_headers") + + # Buffering held every in-flight body twice in router memory (the read + # buffer plus the transport's copy) and serialised the two hops: nothing + # left for the client until the worker finished sending. Chunk through + # instead; the body-side marks go with the buffer, since the header is + # on the wire before the body has passed. + out_headers = dict(response.headers) + for hop in ("content-length", "transfer-encoding", "connection", "date", "server"): + out_headers.pop(hop, None) + out_headers[ROUTER_TIMING_HEADER] = stamps.to_header() + out_headers[ROUTER_WORKER_HEADER] = worker_url + + async def relay(): + try: + async for chunk in response.aiter_raw(): + yield chunk + finally: + await response.aclose() + self._finish_url(worker_url) + + return StreamingResponse(relay(), status_code=response.status_code, headers=out_headers) async def add_worker(self, request: Request): """Add a new worker to the router. diff --git a/miles/utils/diffusion_rollout_response.py b/miles/utils/diffusion_rollout_response.py index 4b1eb6ba..ed82d949 100644 --- a/miles/utils/diffusion_rollout_response.py +++ b/miles/utils/diffusion_rollout_response.py @@ -2,6 +2,7 @@ from __future__ import annotations +import time from collections.abc import Callable from typing import Any @@ -215,6 +216,9 @@ def apply_rollout_image_response( @ray.remote(num_cpus=1) class RolloutImageResponseParserActor: - def apply_raw(self, samples: list[Sample], raw: bytes) -> list[Sample]: + def apply_raw(self, samples: list[Sample], raw: bytes) -> tuple[list[Sample], dict[str, float]]: + """Fill ``samples``, and report when this actor picked the work up so the + caller can split its ray round trip into queueing and real work.""" + marks = {"parser_start": time.time()} bodies = msgpack.unpackb(raw, raw=False) - return [apply_rollout_image_response(s, b) for s, b in zip(samples, bodies, strict=True)] + return [apply_rollout_image_response(s, b) for s, b in zip(samples, bodies, strict=True)], marks diff --git a/miles/utils/http_utils.py b/miles/utils/http_utils.py index 4ea3e740..a6b19822 100644 --- a/miles/utils/http_utils.py +++ b/miles/utils/http_utils.py @@ -5,6 +5,8 @@ import os import random import socket +import time +from typing import Any, NamedTuple import httpx import msgpack @@ -143,11 +145,31 @@ def _parse_response(response, raw=False): return response.text -async def _post(client, url, payload, max_retries=60, raw=False): +class PostResult(NamedTuple): + body: Any + headers: dict[str, str] + marks: dict[str, float] + + +async def _send_once(client, url, payload, timed): + """One POST attempt. Streamed when timed: a rollout reply's first byte and + last byte are seconds apart, and only the sender sees both.""" + if not timed: + return await client.post(url, json=payload or {}), {} + + marks = {"http_send": time.time()} + async with client.stream("POST", url, json=payload or {}) as response: + marks["http_headers"] = time.time() + await response.aread() + marks["http_recv_done"] = time.time() + return response, marks + + +async def _post(client, url, payload, max_retries=60, raw=False, timed=False): retry_count = 0 while retry_count < max_retries: try: - response = await client.post(url, json=payload or {}) + response, marks = await _send_once(client, url, payload, timed) response.raise_for_status() output = _parse_response(response, raw=raw) except Exception as e: @@ -168,6 +190,8 @@ async def _post(client, url, payload, max_retries=60, raw=False): continue break + if timed: + return PostResult(output, dict(response.headers), marks) return output @@ -218,8 +242,8 @@ def __init__(self, concurrency: int): timeout=httpx.Timeout(None), ) - async def do_post(self, url, payload, max_retries=60, raw=False): - return await _post(self._client, url, payload, max_retries, raw) + async def do_post(self, url, payload, max_retries=60, raw=False, timed=False): + return await _post(self._client, url, payload, max_retries, raw, timed) # Create actors per node created = [] @@ -243,7 +267,9 @@ async def do_post(self, url, payload, max_retries=60, raw=False): _post_actors = created -async def post(url, payload, max_retries=60, raw=False): +async def post(url, payload, max_retries=60, raw=False, timed=False): + """POST ``payload``. With ``timed``, return a :class:`PostResult` instead of + just the body, carrying the reply headers and this hop's own timestamps.""" # If distributed mode is enabled and actors exist, dispatch via Ray. if _distributed_post_enabled and _post_actors: try: @@ -252,10 +278,10 @@ async def post(url, payload, max_retries=60, raw=False): actor = _next_actor() if actor is not None: # Use a thread to avoid blocking the event loop on ray.get - obj_ref = actor.do_post.remote(url, payload, max_retries, raw) + obj_ref = actor.do_post.remote(url, payload, max_retries, raw, timed) return await asyncio.to_thread(ray.get, obj_ref) except Exception as e: logger.info(f"[http_utils] Distributed POST failed, falling back to local: {e} (url={url})") # fall through to local - return await _post(_http_client, url, payload, max_retries, raw) + return await _post(_http_client, url, payload, max_retries, raw, timed) diff --git a/miles/utils/request_timing.py b/miles/utils/request_timing.py new file mode 100644 index 00000000..1d35be93 --- /dev/null +++ b/miles/utils/request_timing.py @@ -0,0 +1,182 @@ +"""Timing for one rollout request across the processes it visits. + +Each process stamps absolute marks and returns them on a response header, so +the rollout manager assembles the chain from a reply it already waits for. +Every leg is the interval between two consecutive marks, so a reader places all +of them at the time they happened. Marks from one process are only comparable +within it, so a cross-process leg's own value carries the two clocks' unknown +offset; only the sum of those legs is exact, and ``Timing.cross_total`` reports +it by subtraction. +""" + +from __future__ import annotations + +import json +import time +from dataclasses import dataclass, field + +# Keep in sync with sglang.multimodal_gen.runtime.entrypoints.post_training.timing; +# duplicated rather than imported so an older sglang-diffusion still works. +SGLD_TIMING_HEADER = "x-sgld-timing" +SGLD_STAGES_HEADER = "x-sgld-stages" +ROUTER_TIMING_HEADER = "x-miles-router-timing" +ROUTER_WORKER_HEADER = "x-miles-router-worker" + +# wire key -> mark name; abbreviated because a whole map travels in one header +SGLD_WIRE_KEYS = { + "rc": "srv_recv", + "fs": "forward_start", + "fe": "forward_end", + "bs": "build_start", + "be": "build_end", + "de": "dump_end", + "me": "msgpack_end", +} +ROUTER_WIRE_KEYS = { + "rr": "router_recv", + "rd": "router_dispatch", + "wh": "worker_headers", + "bd": "router_body_done", + "rp": "router_reply", +} + +CLIENT = "client" +ROUTER = "router" +SGLD = "sgld" +CROSS = "cross" # a leg between two processes, hence two clocks + +# Two marks share a clock iff they share a source. The response parser actors +# are pinned to the rollout manager's node, so they read the same host clock. +MARK_SOURCE = { + "req_start": CLIENT, + "slot_acquired": CLIENT, + "http_send": CLIENT, + "http_headers": CLIENT, + "http_recv_done": CLIENT, + "parser_submit": CLIENT, + "parser_start": CLIENT, + "parser_done": CLIENT, + "reward_start": CLIENT, + "reward_end": CLIENT, + "router_recv": ROUTER, + "router_dispatch": ROUTER, + "worker_headers": ROUTER, + "router_body_done": ROUTER, + "router_reply": ROUTER, + "srv_recv": SGLD, + "forward_start": SGLD, + "forward_end": SGLD, + "build_start": SGLD, + "build_end": SGLD, + "dump_end": SGLD, + "msgpack_end": SGLD, +} + +# Consecutive marks, so the legs tile the request and nothing hides in a gap. +# This order is the order every reader renders in. +LEGS: tuple[tuple[str, str, str], ...] = ( + ("wait_slot", "req_start", "slot_acquired"), + ("post_dispatch", "slot_acquired", "http_send"), + ("req_to_router", "http_send", "router_recv"), + ("router_intake", "router_recv", "router_dispatch"), + ("req_to_sgld", "router_dispatch", "srv_recv"), + ("sgld_prepare", "srv_recv", "forward_start"), + ("sgld_forward", "forward_start", "forward_end"), + ("serialize_dispatch", "forward_end", "build_start"), + ("build_response", "build_start", "build_end"), + ("model_dump", "build_end", "dump_end"), + ("msgpack", "dump_end", "msgpack_end"), + ("resp_to_router", "msgpack_end", "worker_headers"), + ("body_sgld_to_router", "worker_headers", "router_body_done"), + ("router_relay", "router_body_done", "router_reply"), + ("resp_to_client", "router_reply", "http_headers"), + ("body_router_to_client", "http_headers", "http_recv_done"), + ("parser_dispatch", "http_recv_done", "parser_submit"), + ("parser_queue", "parser_submit", "parser_start"), + ("parser_work", "parser_start", "parser_done"), + ("reward_wait", "parser_done", "reward_start"), + ("reward", "reward_start", "reward_end"), +) + +LEG_NAMES = tuple(name for name, _, _ in LEGS) + + +def leg_sources() -> dict[str, str]: + """Leg name -> the process that measured it, or ``CROSS`` across two.""" + return {name: _source(a, b) for name, a, b in LEGS} + + +def _source(mark_a: str, mark_b: str) -> str: + src_a, src_b = MARK_SOURCE[mark_a], MARK_SOURCE[mark_b] + return src_a if src_a == src_b else CROSS + + +class Marks: + """Absolute wall-clock marks taken by one process.""" + + __slots__ = ("marks", "_wire") + + def __init__(self, wire_keys: dict[str, str] | None = None) -> None: + self.marks: dict[str, float] = {} + self._wire = {name: key for key, name in (wire_keys or {}).items()} + + def mark(self, name: str) -> None: + assert name in MARK_SOURCE, f"unknown timing mark {name!r}" + self.marks[name] = time.time() + + def absorb(self, marks: dict[str, float]) -> None: + self.marks.update(marks) + + def to_header(self) -> str: + return json.dumps( + {self._wire[name]: round(t, 6) for name, t in self.marks.items() if name in self._wire}, + separators=(",", ":"), + ) + + +def parse_header(value: str | None, wire_keys: dict[str, str]) -> dict[str, float]: + """Decode one timing header; no header means a peer that predates it.""" + if not value: + return {} + raw = json.loads(value) + return {name: float(raw[key]) for key, name in wire_keys.items() if key in raw} + + +def parse_stages(value: str | None) -> dict[str, float]: + """Engine stage seconds, from the milliseconds the header carries.""" + if not value: + return {} + return {name: float(ms) / 1000.0 for name, ms in json.loads(value).items()} + + +@dataclass +class Timing: + durations: dict[str, float] = field(default_factory=dict) + cross_total: float = 0.0 + t_start: float = 0.0 + t_end: float = 0.0 + + +def derive(marks: dict[str, float]) -> Timing: + """Per-leg seconds for one request. + + A cross-process leg's own value carries the two clocks' offset, so a + negative one is how skew shows up. Their sum does not: the marks are + consecutive, so the offsets on the marks between them cancel, leaving + ``cross_total`` as the client's own window minus the in-process legs. + """ + timing = Timing() + in_process = 0.0 + for name, mark_a, mark_b in LEGS: + if mark_a not in marks or mark_b not in marks: + continue + timing.durations[name] = marks[mark_b] - marks[mark_a] + if _source(mark_a, mark_b) != CROSS: + in_process += timing.durations[name] + + client_marks = [t for name, t in marks.items() if MARK_SOURCE[name] == CLIENT] + if client_marks: + timing.t_start = min(client_marks) + timing.t_end = max(client_marks) + timing.cross_total = timing.t_end - timing.t_start - in_process + return timing diff --git a/tests/fast/dashboard/test_async_collector.py b/tests/fast/dashboard/test_async_collector.py index 369816fa..ee5c938a 100644 --- a/tests/fast/dashboard/test_async_collector.py +++ b/tests/fast/dashboard/test_async_collector.py @@ -33,12 +33,14 @@ class FakeHandle: def __init__(self, *, fail: bool = False) -> None: self.push_phases = FakeRemoteMethod(fail=fail) self.push_trajectories = FakeRemoteMethod(fail=fail) + self.push_requests = FakeRemoteMethod(fail=fail) @pytest.fixture(autouse=True) def clean_hook_state(monkeypatch): monkeypatch.setattr(hooks, "_phase_sink", None) monkeypatch.setattr(hooks, "_trajectory_sink", None) + monkeypatch.setattr(hooks, "_request_sink", None) monkeypatch.setattr(hooks, "_GPU_SAMPLER", None) monkeypatch.setattr(hooks, "_resolve_identity", lambda: ("node-a", [2], 7)) monkeypatch.setattr(hooks, "_ray_get", lambda ref: ref) @@ -165,7 +167,8 @@ def test_collector_output_loads_in_existing_viewer(tmp_path): collector.push_trajectories([_trajectory(10.0, "gen_start"), _trajectory(12.0, "gen_end")]) collector.flush() - phases, gpu, lifecycle = load_streams(str(workspace)) + phases, gpu, lifecycle, requests = load_streams(str(workspace)) + assert requests == [] assert phases[0]["name"] == "actor_train" assert gpu == [] assert lifecycle == [{"rollout_id": 5, "sample_index": 3, "stage": "gen", "t0": 10.0, "t1": 12.0}] @@ -176,5 +179,5 @@ def test_collector_shutdown_flushes_tail(tmp_path): collector = DashboardCollector(CollectorConfig(workspace=str(workspace), run_name="test", start_ts=0)) collector.push_phases([_phase(11.0)]) collector.shutdown() - phases, _, _ = load_streams(str(workspace)) + phases, *_ = load_streams(str(workspace)) assert [phase["name"] for phase in phases] == ["actor_train"] diff --git a/tests/fast/dashboard/test_request_timing.py b/tests/fast/dashboard/test_request_timing.py new file mode 100644 index 00000000..a20199a4 --- /dev/null +++ b/tests/fast/dashboard/test_request_timing.py @@ -0,0 +1,178 @@ +"""Per-request rollout timing: mark chain -> leg durations -> dashboard stream. + +One request walks three processes, each stamping its own clock. Every leg is the +interval between two consecutive marks, so the legs tile the request and each +one is drawn at the time it happened -- including the four that cross a process +boundary, which are named individually rather than merged. + + client │ req_start ─wait_slot─ slot_acquired ─── http_send ....... http_recv_done ─parser─ reward_end + │ ╲ ╱ + router │ router_recv ─intake─ router_dispatch ... router_reply + │ ╲ ╱ + sgld │ srv_recv ─forward─ build ─dump─ msgpack_end + ╰── ╲ ╱ = cross-process leg: its own value carries the two clocks' + offset, only the sum of the four is exact (`cross_total`) + +The engine additionally reports its own breakdown of `sgld_forward`, the one leg +it owns and the client cannot see into; those are durations with no marks, so +they are tabulated rather than placed on the axis. + +Covered here: the legs tile exactly (1), `cross_total` survives a clock offset +while the individual hops do not (2), a peer without headers degrades instead of +failing (3), the header round trips, marks and engine stages alike (4-6), and +the record reaching its own stream unaltered (7-8). +""" + +from tests.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=10, suite="stage-a-cpu", labels=[]) + +from types import SimpleNamespace + +import pytest + +from miles.dashboard import hooks +from miles.dashboard.collector import CollectorConfig, DashboardCollector +from miles.dashboard.hooks import RequestSink, RequestTracer +from miles.dashboard.store import Stream +from miles.dashboard.viewer import load_streams +from miles.utils.request_timing import ( + CROSS, + LEGS, + MARK_SOURCE, + ROUTER_WIRE_KEYS, + Marks, + derive, + leg_sources, + parse_header, + parse_stages, +) + +# one plausible request as a gap per leg; the cross-process gaps are the small +# ones, since the large bodies move inside a span one process measures alone +GAPS = { + "wait_slot": 1.5, + "post_dispatch": 0.001, + "req_to_router": 0.002, + "router_intake": 0.001, + "req_to_sgld": 0.004, + "sgld_prepare": 0.003, + "sgld_forward": 11.0, + "serialize_dispatch": 0.0002, + "build_response": 1.0, + "model_dump": 0.1, + "msgpack": 0.5, + "resp_to_router": 0.02, + "body_sgld_to_router": 2.0, + "router_relay": 0.001, + "resp_to_client": 0.01, + "body_router_to_client": 1.5, + "parser_dispatch": 0.0004, + "parser_queue": 0.3, + "parser_work": 0.8, + "reward_wait": 0.001, + "reward": 0.4, +} +LEG_SOURCE = leg_sources() +CROSS_LEGS = [name for name in GAPS if LEG_SOURCE[name] == CROSS] +CROSS_TOTAL = sum(GAPS[name] for name in CROSS_LEGS) + + +def _marks(*, offsets=None, drop_sources=()): + """Walk the chain, adding each process's clock offset to its own marks.""" + offsets = offsets or {} + marks, t = {}, 1_000_000.0 + for name, mark_a, mark_b in LEGS: + marks.setdefault(mark_a, t + offsets.get(MARK_SOURCE[mark_a], 0.0)) + t += GAPS[name] + marks[mark_b] = t + offsets.get(MARK_SOURCE[mark_b], 0.0) + return {name: ts for name, ts in marks.items() if MARK_SOURCE[name] not in drop_sources} + + +def _samples(count=2): + return [SimpleNamespace(index=i, group_index=2, request_id="rid-9") for i in range(count)] + + +def _tracer(**marks): + tracer = RequestTracer(rollout_id=4) + tracer.marks.absorb(marks or _marks()) + return tracer + + +def test_every_leg_is_measured_and_they_tile_the_request(): + timing = derive(_marks()) + for name, gap in GAPS.items(): + assert timing.durations[name] == pytest.approx(gap, abs=1e-6), name + assert timing.cross_total == pytest.approx(CROSS_TOTAL, abs=1e-6) + assert sum(timing.durations.values()) == pytest.approx(timing.t_end - timing.t_start, abs=1e-6) + + +def test_a_clock_offset_distorts_each_hop_but_not_their_sum(): + timing = derive(_marks(offsets={"router": 5.0, "sgld": -3.0})) + assert timing.cross_total == pytest.approx(CROSS_TOTAL, abs=1e-6) + for name in GAPS: + if LEG_SOURCE[name] != CROSS: + assert timing.durations[name] == pytest.approx(GAPS[name], abs=1e-6), name + assert timing.durations["req_to_sgld"] < 0 # how the offset shows up + + +def test_a_peer_without_timing_headers_degrades_to_client_legs(): + timing = derive(_marks(drop_sources={"router", "sgld"})) + assert set(timing.durations) == {name for name in GAPS if LEG_SOURCE[name] == "client"} + in_process = sum(timing.durations.values()) + assert timing.cross_total == pytest.approx(timing.t_end - timing.t_start - in_process, abs=1e-6) + + +def test_marks_round_trip_through_a_header(): + marks = Marks(ROUTER_WIRE_KEYS) + for name in ROUTER_WIRE_KEYS.values(): + marks.mark(name) + decoded = parse_header(marks.to_header(), ROUTER_WIRE_KEYS) + assert decoded.keys() == marks.marks.keys() + for name, value in decoded.items(): + assert value == pytest.approx(marks.marks[name], abs=1e-6) + + +def test_an_absent_header_yields_no_marks(): + assert parse_header(None, ROUTER_WIRE_KEYS) == {} + + +def test_engine_stages_arrive_as_seconds(): + # the engine reports milliseconds; everything downstream is seconds + assert parse_stages('{"decoding":8123.5,"text_encoding":91.2}') == pytest.approx( + {"decoding": 8.1235, "text_encoding": 0.0912}, abs=1e-9 + ) + assert parse_stages(None) == {} + + +def test_sink_records_the_marks_one_event_per_request(): + handle = SimpleNamespace(push_requests=SimpleNamespace(remote=lambda batch: batch)) + sink = RequestSink(handle) + tracer = _tracer() + tracer.worker = "http://worker:30000" + tracer.resp_bytes = 7 + tracer.engine_stages = {"decoding": 8.1} + + sink.record(tracer, _samples()) + (event,) = sink._buffer + + assert (event.rollout_id, event.request_id, event.group_index) == (4, "rid-9", 2) + assert event.sample_indices == [0, 1] + assert event.marks == pytest.approx(tracer.marks.marks, abs=1e-6) + assert (event.worker, event.resp_bytes) == ("http://worker:30000", 7) + assert event.engine_stages == {"decoding": 8.1} + + +def test_request_events_reach_their_own_stream(tmp_path, monkeypatch): + monkeypatch.setattr(hooks, "_ray_get", lambda ref: ref) + workspace = tmp_path / "dashboard" + collector = DashboardCollector(CollectorConfig(workspace=str(workspace), run_name="test", start_ts=0)) + sink = RequestSink(SimpleNamespace(push_requests=SimpleNamespace(remote=collector.push_requests))) + sink.record(_tracer(), _samples()) + sink.flush() + collector.flush() + + assert list((workspace / Stream.REQUESTS.value).glob("*.jsonl")) + _, _, _, requests = load_streams(str(workspace)) + assert [r["request_id"] for r in requests] == ["rid-9"] + assert derive(requests[0]["marks"]).durations["sgld_forward"] == pytest.approx(GAPS["sgld_forward"], abs=1e-4)