Skip to content
39 changes: 11 additions & 28 deletions .github/workflows/prek.yml
Original file line number Diff line number Diff line change
Expand Up @@ -51,18 +51,9 @@ jobs:
run: |
fingerprint="${{ steps.fingerprint.outputs.fingerprint }}"
part_prefix="${CI_UV_CACHE_ASSET_PREFIX}-${fingerprint}.tar.zst.part-"
release_api="https://api.github.com/repos/${GITHUB_REPOSITORY}/releases/tags/${CI_UV_CACHE_RELEASE_TAG}"

release_json="$(curl -fsSL \
-H "Authorization: Bearer ${GITHUB_TOKEN}" \
-H "Accept: application/vnd.github+json" \
"${release_api}" || true)"

if [ -z "${release_json}" ]; then
echo "Cache release '${CI_UV_CACHE_RELEASE_TAG}' not found."
echo "cache-hit=false" >> "${GITHUB_OUTPUT}"
exit 0
fi
release_json="$(python3 scripts/ci/github_release_assets.py \
--repository "${GITHUB_REPOSITORY}" \
--tag "${CI_UV_CACHE_RELEASE_TAG}")"

hit="$(RELEASE_JSON="${release_json}" PART_PREFIX="${part_prefix}" python3 -c "
import json, os, re
Expand All @@ -73,7 +64,7 @@ jobs:
int(m.group(1))
for a in payload.get('assets', [])
for m in [pattern.match(a.get('name', ''))]
if m and a.get('id') is not None
if m and a.get('url')
)
print('true' if parts and parts == list(range(len(parts))) else 'false')
")"
Expand Down Expand Up @@ -147,22 +138,15 @@ jobs:
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
run: |
release_api="https://api.github.com/repos/${GITHUB_REPOSITORY}/releases/tags/${CI_UV_CACHE_RELEASE_TAG}"
fingerprint="${{ needs.cache-status.outputs.fingerprint }}"
part_prefix="${CI_UV_CACHE_ASSET_PREFIX}-${fingerprint}.tar.zst.part-"

release_json="$(curl -fsSL \
-H "Authorization: Bearer ${GITHUB_TOKEN}" \
-H "Accept: application/vnd.github+json" \
"${release_api}" || true)"

if [ -z "${release_json}" ]; then
echo "::error::Missing cache release '${CI_UV_CACHE_RELEASE_TAG}'."
exit 1
fi
release_json="$(python3 scripts/ci/github_release_assets.py \
--repository "${GITHUB_REPOSITORY}" \
--tag "${CI_UV_CACHE_RELEASE_TAG}")"

part_selection_file="/tmp/uv-cache-part-selection.txt"
if ! RELEASE_JSON="${release_json}" PART_PREFIX="${part_prefix}" python3 -c "import json, os, re, sys; payload=json.loads(os.environ['RELEASE_JSON']); part_prefix=os.environ['PART_PREFIX']; pattern=re.compile(r'^' + re.escape(part_prefix) + r'(\\d{3})$'); parts=[]; [parts.append((int(m.group(1)), int(a.get('id')), a.get('name'))) for a in payload.get('assets', []) for m in [pattern.match(a.get('name', ''))] if m and a.get('id') is not None]; parts.sort(key=lambda x: x[0]); indices=[p[0] for p in parts]; expected=list(range(len(parts))); print('\\n'.join(f'{asset_id} {name}' for _, asset_id, name in parts)) if parts and indices == expected else (_ for _ in ()).throw(SystemExit(2 if not parts else 3))" > "${part_selection_file}"; then
if ! RELEASE_JSON="${release_json}" PART_PREFIX="${part_prefix}" python3 -c "import json, os, re, sys; payload=json.loads(os.environ['RELEASE_JSON']); part_prefix=os.environ['PART_PREFIX']; pattern=re.compile(r'^' + re.escape(part_prefix) + r'(\\d{3})$'); parts=[]; [parts.append((int(m.group(1)), a.get('url'), a.get('name'))) for a in payload.get('assets', []) for m in [pattern.match(a.get('name', ''))] if m and a.get('url')]; parts.sort(key=lambda x: x[0]); indices=[p[0] for p in parts]; expected=list(range(len(parts))); print('\\n'.join(f'{url} {name}' for _, url, name in parts)) if parts and indices == expected else (_ for _ in ()).throw(SystemExit(2 if not parts else 3))" > "${part_selection_file}"; then
echo "::error::No complete uv cache part set found for prefix '${part_prefix}'."
exit 1
fi
Expand All @@ -176,15 +160,14 @@ jobs:
mkdir -p "${parts_dir}"
awk -v d="${parts_dir}" '{print d "/" $2}' "${part_selection_file}" > "${part_paths_file}"

PARTS_DIR="${parts_dir}" GITHUB_TOKEN="${GITHUB_TOKEN}" GITHUB_REPOSITORY="${GITHUB_REPOSITORY}" \
PARTS_DIR="${parts_dir}" GITHUB_TOKEN="${GITHUB_TOKEN}" \
xargs -n 2 -P 8 sh -c '
asset_id="$1"
asset_url="$1"
asset_name="$2"
part_path="${PARTS_DIR}/${asset_name}"
curl -fsSL -L \
-H "Authorization: Bearer ${GITHUB_TOKEN}" \
-H "Accept: application/octet-stream" \
"https://api.github.com/repos/${GITHUB_REPOSITORY}/releases/assets/${asset_id}" \
"${asset_url}" \
-o "${part_path}"
' sh < "${part_selection_file}"

Expand Down
58 changes: 48 additions & 10 deletions dev/trainer_rank_check.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,7 @@ def main(
depths: str = "0,1,2,3,4",
performance_depth: int = 1,
chunks: str = "17,512,8192",
workload: Literal["regular", "austin", "varied"] = "regular",
workload: Literal["regular", "austin", "varied", "unequal_slots"] = "regular",
request: Literal["target", "multi", "topk", "logits", "hidden", "mixed"] = "target",
families: int = 8,
prefix_tokens: int = 128,
Expand Down Expand Up @@ -77,6 +77,8 @@ def main(
),
print_env=dist.get_rank() == 0,
)
if mode == "performance" and workload == "unequal_slots":
_trace_unequal_slots("runtime_ready")
for chunk in runtime.model:
chunk.eval()
if mode == "correctness":
Expand Down Expand Up @@ -396,21 +398,46 @@ def _performance(
if workload == "austin":
families, prefix_tokens, branches, completion_tokens = 30, 5000, 16, 100
rank = TrainerRank(runtime, shared_prefix_max_depth=depth, head_chunk_tokens=8192)
slot_names = load_random_checkpoint_slots(runtime, rank, slots)
requests = _performance_requests(
request=request,
families=families,
prefix_tokens=prefix_tokens,
branches=branches,
completion_tokens=completion_tokens,
varied=workload == "varied",
slots=slot_names,
if workload == "unequal_slots":
_trace_unequal_slots("rank_ready")
slot_names = load_random_checkpoint_slots(
runtime,
rank,
slots,
site_limit=1 if workload == "unequal_slots" else None,
)
if workload == "unequal_slots":
_trace_unequal_slots("slots_ready")
if workload == "unequal_slots":
if len(slot_names) < 2:
raise ValueError("--workload unequal_slots requires --slots >= 2")
requests = [
ForwardInput(
input_tokens=(tokens := _tokens(index * 1009, length)),
target_tokens=(tokens * 7 + 3) % 32_000,
checkpoint=slot_names[index],
)
for index, length in enumerate((256, 64))
]
else:
requests = _performance_requests(
request=request,
families=families,
prefix_tokens=prefix_tokens,
branches=branches,
completion_tokens=completion_tokens,
varied=workload == "varied",
slots=slot_names,
)
dp_rank, dp_size = rank._dp_rank_and_size()
plan = rank._plan_flat_forward(requests)
if workload == "unequal_slots":
_trace_unequal_slots("plan_ready")
assert workload != "austin" or plan.packed_tokens == 198_000

def step() -> list[MicroBatchStats]:
if workload == "unequal_slots":
_trace_unequal_slots("step_start")
rank.zero_grad()
stats: list[MicroBatchStats] = []
if adaptive:
Expand All @@ -419,7 +446,11 @@ def step() -> list[MicroBatchStats]:
stats.append(micro.stats)
else:
outputs = rank.dp_rank_forward(requests[dp_rank::dp_size])
if workload == "unequal_slots":
_trace_unequal_slots("forward_ready")
_output_loss(outputs).backward()
if workload == "unequal_slots":
_trace_unequal_slots("backward_ready")
if optimizer_step:
if not slot_names:
raise ValueError("--optimizer-step requires --slots >= 1")
Expand Down Expand Up @@ -461,6 +492,13 @@ def step() -> list[MicroBatchStats]:
}


def _trace_unequal_slots(event: str) -> None:
print(
json.dumps({"event": event, "rank": dist.get_rank(), "time": time.time()}),
flush=True,
)


def _output_loss(outputs: Iterable[ForwardOutput]) -> torch.Tensor:
terms: list[torch.Tensor] = []
for output in outputs:
Expand Down
15 changes: 14 additions & 1 deletion dev/trainer_rank_support.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ def load_random_checkpoint_slots(
count: int,
*,
lora_rank: int = 8,
site_limit: int | None = None,
) -> tuple[str, ...]:
assert count >= 0, "slots must be >= 0"
if count == 0:
Expand All @@ -23,12 +24,24 @@ def load_random_checkpoint_slots(
gathered, LoRAPublishPlanner(runtime.model).global_metadata({})
)
metadata = {meta.key: meta for values in gathered if values for meta in values}
selected = sorted(metadata.values(), key=lambda item: item.key)
if site_limit is not None:
pairs = []
for meta in selected:
if ".lora_A." not in meta.key or ".experts." in meta.key:
continue
b_key = meta.key.replace(".lora_A.", ".lora_B.")
if b_meta := metadata.get(b_key):
pairs.append((meta, b_meta))
selected = [meta for pair in pairs[:site_limit] for meta in pair]
if not selected:
raise RuntimeError("No replicated LoRA sites are available for the check")
dtype = next(runtime.model[0].parameters()).dtype
names = tuple(f"S{index}" for index in range(count))
for index, name in enumerate(names):
generator = torch.Generator(device=rank.device).manual_seed(index + 1)
adapter: dict[str, torch.Tensor] = {}
for meta in sorted(metadata.values(), key=lambda item: item.key):
for meta in selected:
shape = list(meta.shape)
if meta.manifest["sharded"]:
axis = int(meta.manifest["export_shard_dim"])
Expand Down
78 changes: 78 additions & 0 deletions scripts/ci/github_release_assets.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,78 @@
#!/usr/bin/env python3
"""List GitHub release assets without the REST release endpoint."""

from __future__ import annotations

import argparse
import json
import os
from urllib import request

_QUERY = """
query($owner: String!, $name: String!, $tag: String!, $cursor: String) {
repository(owner: $owner, name: $name) {
release(tagName: $tag) {
releaseAssets(first: 100, after: $cursor) {
nodes { name downloadUrl }
pageInfo { hasNextPage endCursor }
}
}
}
}
"""


def release_assets(repository: str, tag: str, token: str) -> list[dict[str, str]]:
owner, name = repository.split("/", 1)
cursor: str | None = None
assets: list[dict[str, str]] = []
while True:
body = json.dumps(
{
"query": _QUERY,
"variables": {
"owner": owner,
"name": name,
"tag": tag,
"cursor": cursor,
},
}
).encode()
http_request = request.Request(
"https://api.github.com/graphql",
data=body,
headers={
"Accept": "application/vnd.github+json",
"Authorization": f"Bearer {token}",
"Content-Type": "application/json",
},
)
with request.urlopen(http_request, timeout=30) as response:
payload = json.load(response)
if errors := payload.get("errors"):
raise RuntimeError(f"GitHub GraphQL error: {errors}")
release = payload["data"]["repository"]["release"]
if release is None:
return []
page = release["releaseAssets"]
assets.extend(
{"name": node["name"], "url": node["downloadUrl"]} for node in page["nodes"]
)
if not page["pageInfo"]["hasNextPage"]:
return assets
cursor = page["pageInfo"]["endCursor"]


def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--repository", required=True)
parser.add_argument("--tag", required=True)
args = parser.parse_args()
token = os.environ.get("GITHUB_TOKEN")
if not token:
raise SystemExit("GITHUB_TOKEN is required")
print(json.dumps({"assets": release_assets(args.repository, args.tag, token)}))


if __name__ == "__main__":
main()
Loading
Loading