feat: add weekly workflow to update CI test est_time values (#22545)
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.6
parent
f7a1740101
commit
3f39b3d811
@@ -0,0 +1,70 @@
|
||||
name: Weekly Update Est Time
|
||||
|
||||
on:
|
||||
schedule:
|
||||
- cron: '0 0 * * 6' # Saturday 00:00 UTC
|
||||
workflow_dispatch: {}
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
pull-requests: write
|
||||
|
||||
jobs:
|
||||
update-est-time:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
token: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.10'
|
||||
|
||||
- name: Update est_time values
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
run: |
|
||||
python scripts/ci/update_est_time.py
|
||||
|
||||
- name: Check for changes
|
||||
id: changes
|
||||
run: |
|
||||
if git diff --quiet; then
|
||||
echo "has_changes=false" >> "$GITHUB_OUTPUT"
|
||||
echo "No est_time changes detected"
|
||||
else
|
||||
echo "has_changes=true" >> "$GITHUB_OUTPUT"
|
||||
echo "Est_time changes detected:"
|
||||
git diff --stat
|
||||
fi
|
||||
|
||||
- name: Create PR
|
||||
if: steps.changes.outputs.has_changes == 'true'
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GH_PAT_FOR_PULL_REQUEST }}
|
||||
run: |
|
||||
git config user.name "sglang-bot"
|
||||
git config user.email "sglang-bot@users.noreply.github.com"
|
||||
|
||||
BRANCH_NAME="bot/update-est-time-$(date +%Y%m%d)"
|
||||
git checkout -b "$BRANCH_NAME"
|
||||
|
||||
git add -A
|
||||
git commit -m "chore: update CI test est_time from recent run data"
|
||||
|
||||
git push origin "$BRANCH_NAME"
|
||||
|
||||
gh pr create \
|
||||
--title "chore: update CI test est_time values" \
|
||||
--body "## Summary
|
||||
|
||||
Updates \`est_time\` values in CI test registration calls based on the median of the last 10 successful executions from scheduled PR Test runs on main.
|
||||
|
||||
This keeps the LPT load-balancing algorithm accurate for partitioning tests across parallel CI jobs.
|
||||
|
||||
🤖 Generated with GitHub Actions" \
|
||||
--base main \
|
||||
--head "$BRANCH_NAME"
|
||||
Executable
+253
@@ -0,0 +1,253 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Update est_time values in CI test files based on actual execution times.
|
||||
|
||||
Fetches logs from recent scheduled PR Test workflow runs on main,
|
||||
parses per-file elapsed times from successful jobs, computes medians,
|
||||
and updates the est_time literals in test registration calls.
|
||||
|
||||
Usage:
|
||||
python scripts/ci/update_est_time.py [--dry-run] [--repo OWNER/REPO]
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import re
|
||||
import statistics
|
||||
import subprocess
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parent.parent.parent
|
||||
|
||||
# Regex to extract per-file elapsed time from CI logs.
|
||||
# Matches lines like:
|
||||
# filename='/actions-runner/_work/sglang/sglang/test/registered/core/test_x.py', elapsed=120, ...
|
||||
# filename='/actions-runner/_work/sglang/sglang/python/sglang/jit_kernel/tests/test_x.py', elapsed=120, ...
|
||||
LOG_PATTERN = re.compile(
|
||||
r"filename='[^']*?/sglang/((?:test|python)/[^']+\.py)', elapsed=(\d+),"
|
||||
)
|
||||
|
||||
WORKFLOW_NAME = "PR Test"
|
||||
MIN_DATA_POINTS = 3
|
||||
TARGET_DATA_POINTS = 10
|
||||
MAX_RUNS = 20
|
||||
|
||||
|
||||
def gh_api(endpoint, paginate=False):
|
||||
"""Call gh api and return parsed JSON."""
|
||||
cmd = ["gh", "api", endpoint]
|
||||
if paginate:
|
||||
cmd.append("--paginate")
|
||||
result = subprocess.run(cmd, capture_output=True, text=True, check=True)
|
||||
return json.loads(result.stdout)
|
||||
|
||||
|
||||
def gh_api_raw(endpoint):
|
||||
"""Call gh api and return raw bytes (for log downloads)."""
|
||||
cmd = ["gh", "api", endpoint]
|
||||
result = subprocess.run(cmd, capture_output=True, check=True)
|
||||
return result.stdout
|
||||
|
||||
|
||||
def get_workflow_id(repo):
|
||||
"""Find the workflow ID for the PR Test workflow."""
|
||||
data = gh_api(f"/repos/{repo}/actions/workflows")
|
||||
for wf in data["workflows"]:
|
||||
if wf["name"] == WORKFLOW_NAME:
|
||||
return wf["id"]
|
||||
raise RuntimeError(f"Workflow '{WORKFLOW_NAME}' not found in {repo}")
|
||||
|
||||
|
||||
def get_scheduled_runs(repo, workflow_id):
|
||||
"""Get completed scheduled runs on main, newest first."""
|
||||
data = gh_api(
|
||||
f"/repos/{repo}/actions/workflows/{workflow_id}/runs"
|
||||
f"?branch=main&status=completed&event=schedule&per_page=100"
|
||||
)
|
||||
return data["workflow_runs"]
|
||||
|
||||
|
||||
def get_successful_jobs(repo, run_id):
|
||||
"""Get successful jobs for a given run."""
|
||||
data = gh_api(f"/repos/{repo}/actions/runs/{run_id}/jobs?per_page=100")
|
||||
return [j for j in data["jobs"] if j["conclusion"] == "success"]
|
||||
|
||||
|
||||
def determine_backend(job_name):
|
||||
"""Determine backend from job name."""
|
||||
name = job_name.lower()
|
||||
for backend in ["cpu", "amd", "npu"]:
|
||||
if backend in name:
|
||||
return backend
|
||||
return "cuda"
|
||||
|
||||
|
||||
def parse_job_logs(repo, job_id):
|
||||
"""Download and parse a job's logs for elapsed times.
|
||||
|
||||
Returns list of (relative_path, elapsed_seconds) tuples.
|
||||
"""
|
||||
try:
|
||||
raw = gh_api_raw(f"/repos/{repo}/actions/jobs/{job_id}/logs")
|
||||
text = raw.decode("utf-8", errors="replace")
|
||||
except subprocess.CalledProcessError:
|
||||
return []
|
||||
|
||||
results = []
|
||||
for match in LOG_PATTERN.finditer(text):
|
||||
rel_path = match.group(1)
|
||||
elapsed = int(match.group(2))
|
||||
results.append((rel_path, elapsed))
|
||||
return results
|
||||
|
||||
|
||||
def collect_timings(repo):
|
||||
"""Collect per-file elapsed times from recent scheduled CI runs.
|
||||
|
||||
Returns dict mapping (relative_path, backend) -> list of elapsed times
|
||||
(newest first).
|
||||
"""
|
||||
workflow_id = get_workflow_id(repo)
|
||||
print(f"Found workflow '{WORKFLOW_NAME}' (id={workflow_id})")
|
||||
|
||||
runs = get_scheduled_runs(repo, workflow_id)
|
||||
print(f"Found {len(runs)} completed scheduled runs on main")
|
||||
|
||||
# timings[(rel_path, backend)] = [elapsed1, elapsed2, ...]
|
||||
timings = defaultdict(list)
|
||||
runs_processed = 0
|
||||
|
||||
for run in runs:
|
||||
run_id = run["id"]
|
||||
jobs = get_successful_jobs(repo, run_id)
|
||||
if not jobs:
|
||||
continue
|
||||
|
||||
runs_processed += 1
|
||||
test_jobs = [
|
||||
j
|
||||
for j in jobs
|
||||
if j["name"] != "check-changes" and "health" not in j["name"].lower()
|
||||
]
|
||||
print(
|
||||
f" Run {run_id} ({run['conclusion']}): "
|
||||
f"{len(test_jobs)} successful test jobs"
|
||||
)
|
||||
|
||||
for job in test_jobs:
|
||||
backend = determine_backend(job["name"])
|
||||
entries = parse_job_logs(repo, job["id"])
|
||||
for rel_path, elapsed in entries:
|
||||
key = (rel_path, backend)
|
||||
timings[key].append(elapsed)
|
||||
|
||||
if runs_processed >= MAX_RUNS:
|
||||
print(f" Reached max {MAX_RUNS} runs, stopping collection")
|
||||
break
|
||||
|
||||
print(
|
||||
f"\nProcessed {runs_processed} runs, "
|
||||
f"collected timings for {len(timings)} (file, backend) pairs"
|
||||
)
|
||||
return timings
|
||||
|
||||
|
||||
def compute_medians(timings):
|
||||
"""Compute median of last TARGET_DATA_POINTS timings for each (file, backend).
|
||||
|
||||
Returns dict mapping (rel_path, backend) -> median (int).
|
||||
Only includes entries with >= MIN_DATA_POINTS data points.
|
||||
"""
|
||||
medians = {}
|
||||
for key, values in timings.items():
|
||||
recent = values[:TARGET_DATA_POINTS]
|
||||
if len(recent) < MIN_DATA_POINTS:
|
||||
continue
|
||||
medians[key] = round(statistics.median(recent))
|
||||
return medians
|
||||
|
||||
|
||||
def update_est_times(medians, dry_run=False):
|
||||
"""Update est_time values in source files.
|
||||
|
||||
Returns (updated_count, skipped_count).
|
||||
"""
|
||||
updated = 0
|
||||
skipped = 0
|
||||
|
||||
# Group medians by file
|
||||
by_file = defaultdict(dict)
|
||||
for (rel_path, backend), median in medians.items():
|
||||
by_file[rel_path][backend] = median
|
||||
|
||||
for rel_path, backend_medians in sorted(by_file.items()):
|
||||
filepath = REPO_ROOT / rel_path
|
||||
if not filepath.exists():
|
||||
print(f" SKIP {rel_path}: file not found")
|
||||
skipped += 1
|
||||
continue
|
||||
|
||||
content = filepath.read_text()
|
||||
new_content = content
|
||||
|
||||
for backend, median in backend_medians.items():
|
||||
pattern = re.compile(rf"(register_{backend}_ci\(est_time=)(\d+)")
|
||||
match = pattern.search(new_content)
|
||||
if not match:
|
||||
continue
|
||||
|
||||
old_val = int(match.group(2))
|
||||
if old_val == median:
|
||||
continue
|
||||
|
||||
new_content = pattern.sub(rf"\g<1>{median}", new_content)
|
||||
print(
|
||||
f" {rel_path}: register_{backend}_ci "
|
||||
f"est_time={old_val} -> {median}"
|
||||
)
|
||||
|
||||
if new_content != content:
|
||||
if not dry_run:
|
||||
filepath.write_text(new_content)
|
||||
updated += 1
|
||||
else:
|
||||
skipped += 1
|
||||
|
||||
return updated, skipped
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Update est_time values from CI run data"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dry-run",
|
||||
action="store_true",
|
||||
help="Print changes without modifying files",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--repo",
|
||||
default="sgl-project/sglang",
|
||||
help="GitHub repository (default: sgl-project/sglang)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
print("Collecting timings from CI logs...")
|
||||
timings = collect_timings(args.repo)
|
||||
|
||||
print("\nComputing medians...")
|
||||
medians = compute_medians(timings)
|
||||
print(f"Computed medians for {len(medians)} (file, backend) pairs")
|
||||
|
||||
print("\nUpdating est_time values...")
|
||||
updated, skipped = update_est_times(medians, dry_run=args.dry_run)
|
||||
|
||||
action = "Would update" if args.dry_run else "Updated"
|
||||
print(f"\n{action} {updated} files, skipped {skipped} files")
|
||||
|
||||
if args.dry_run:
|
||||
print("(dry-run mode, no files modified)")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user