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")
| 27 | ] |
| 28 | |
| 29 | |
| 30 | def 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 |
no test coverage detected