MCPcopy Create free account
hub / github.com/FlashSampling/FlashSampling / collect_long_df

Function collect_long_df

benchmarking/plot_tp_scaling.py:29–74  ·  view source on GitHub ↗

Read the per-TP batch-scaling CSVs and return a long DataFrame. When ``args.use_reruns`` is set, also load sibling ``{gpu}-rerun*`` folders so each (tp, n_hidden_states, provider) cell has multiple rows. Slow-pod runs are dropped per ``args.max_slowdown`` (FMMS-Triton median must be

(args: "Args")

Source from the content-addressed store, hash-verified

27]
28
29
30def collect_long_df(args: "Args") -> pd.DataFrame:
31 """Read the per-TP batch-scaling CSVs and return a long DataFrame.
32
33 When ``args.use_reruns`` is set, also load sibling ``{gpu}-rerun*``
34 folders so each (tp, n_hidden_states, provider) cell has multiple rows.
35 Slow-pod runs are dropped per ``args.max_slowdown`` (FMMS-Triton median
36 must be within that factor of the fastest run at the same tp).
37
38 Columns: tp, n_hidden_states, provider, time[ms], run.
39 """
40 fmms = FLASHSAMPLING_RENAMES[L.fmms_triton]
41 parent = args.base_dir / "triton-bench" / args.bench_fn
42 frames: list[pd.DataFrame] = []
43 for tp in args.tps:
44 run_dirs = [parent / args.gpu]
45 if args.use_reruns:
46 run_dirs += sorted(parent.glob(f"{args.gpu}-rerun*"))
47
48 # First pass: read all and compute FMMS median per run
49 per_run: list[tuple[str, float, pd.DataFrame]] = []
50 for run_dir in run_dirs:
51 csv = run_dir / f"tp{tp}" / f"fused-mm-sample-batch-scaling-{args.case}.csv"
52 if not csv.exists():
53 if run_dir.name == args.gpu:
54 raise FileNotFoundError(csv)
55 print(f"warn: missing {csv}")
56 continue
57 wide = read_triton_bench_csv(csv).rename(columns=FLASHSAMPLING_RENAMES)
58 fmms_med = wide[fmms].median()
59 per_run.append((run_dir.name, fmms_med, wide))
60
61 if not per_run:
62 continue
63 fastest = min(m for _, m, _ in per_run)
64 for name, fmms_med, wide in per_run:
65 slowdown = fmms_med / fastest
66 if slowdown > args.max_slowdown:
67 print(f"drop tp{tp} {name}: FMMS={fmms_med * 1000:.1f}us, {slowdown:.2f}x slower")
68 continue
69 long = wide.melt(
70 id_vars=["n_hidden_states"], var_name="provider", value_name="time[ms]"
71 )
72 long["tp"] = tp
73 long["run"] = name
74 frames.append(long)
75 return pd.concat(frames, ignore_index=True)
76
77

Callers 1

plot_tp_scaling.pyFile · 0.85

Calls 1

read_triton_bench_csvFunction · 0.90

Tested by

no test coverage detected