Skip to content

Latest commit

 

History

History
403 lines (327 loc) · 23.2 KB

File metadata and controls

403 lines (327 loc) · 23.2 KB

CLIF-STITCH — Workflow & Functions

How the lead and local sites actually run, expressed as a handful of functions over the storage mirror. Read DESIGN.md first for the folder layout and config model.

Scope. STITCH is the communication protocol + file sync + logging. It does NOT build, train, aggregate, or even read a model — the lead's own train / aggregate scripts do that. STITCH moves two files per pass (model.<ext> + metrics.json) between the local mirror and the blob, hashes the bytes, and returns a Status. The user writes/reads the model file at the paths STITCH hands it.

1. The status codes, the path helpers, and the transfer functions

Every transfer returns one Status; STITCH logs it and the loops branch on it:

class Status(Enum):
    UPLOADED       # push succeeded
    DOWNLOADED     # pull fetched a new pass
    NO_NEW_PASS    # pull: nothing new yet -> poll again
    FAILED         # gave up after transfer.max_retries
    RUN_COMPLETE   # global/DONE.json present -> stop the loop
    ALREADY_SYNCED # local mirror already has it -> skipped (idempotent re-run)

The user never hardcodes a path — STITCH owns the layout and hands out paths:

client.model_path(pass_id)        # ./stitch_store/<run_id>/sites/<id>/pass_NNN/model.<ext>
client.metrics_path(pass_id)      # ...               /sites/<id>/pass_NNN/metrics.json
client.global_model_path(pass_id) # ...               /global/pass_NNN/model.<ext>   (to READ)

The transfer functions take only a pass id (and, for pushes, the metrics dict). They never receive or return a model object — STITCH copies the file at the helper path:

# ---- used by ALL sites (lead included) ----

def pull_global(client, pass_id: int) -> Status:
    """Sync global/pass_{pass_id:03d}/ from blob -> local mirror, JIT.
    Checks global/DONE.json first (-> RUN_COMPLETE). If the pass isn't on the blob yet
    -> NO_NEW_PASS. If already in the local mirror -> ALREADY_SYNCED. Else download both
    files, re-hash model bytes vs metrics.model_sha256 -> DOWNLOADED (or FAILED after retries).
    The USER then opens client.global_model_path(pass_id) and loads it however it likes."""

def push_local_model(client, pass_id: int, metrics: dict) -> Status:
    """The USER has already written its model file to client.model_path(pass_id).
    STITCH hashes that file, writes metrics.json (fixed envelope; metrics MUST contain
    n_samples) LAST as the marker, and uploads the pass folder local -> blob.
    Whitelist + max_file_mb checked. -> UPLOADED | FAILED."""


# ---- used by the LEAD's aggregation script only ----

def pull_sites(client, pass_id: int) -> list[Status]:
    """Sync every sites/*/pass_{pass_id:03d}/ from blob -> local mirror so the user's
    aggregate() can read them. Lists sites/*/pass_NNN/metrics.json to see who finished."""

def push_global_model(client, pass_id: int, metrics: dict) -> Status:
    """The USER has written the aggregated model to client.model_path(pass_id) in the
    global lane. STITCH writes metrics.json, uploads global/pass_NNN/ -> blob, and advances
    orchestration.json. LEAD is the only writer of global/. -> UPLOADED | FAILED."""

pull_global + push_local_model are the only two a local site touches. The lead adds pull_sites + push_global_model, which its own aggregation script calls — STITCH puts the files in the local mirror; the averaging math (and the model loading) is the user's.

Orchestration & event helpers (the instruction set + the insertion folder)

Two more thin helpers, used by the loops, not the user's model code:

client.read_orchestration() -> dict   # the lead's instruction set: current_pass, phase,
                                      #   directive.expected, site_status (read it FIRST)
client.emit(pass_id, event, **fields) # append sites/<id>/events/pass_NNN.<event>.json
                                      #   event in {started, completed, failed}; idempotent key
client.heartbeat(pass_id)             # overwrite sites/<id>/events/heartbeat.json (liveness)
client.resume_point() -> int          # highest pass I've completed (from mirror + events) + 1

emit("completed", ...) is the "report back" — it records sha256 + n_samples + metrics so the lead's site_status rollup (and clif-stitch status) can see the site finished. Only the lead writes orchestration.json; each site writes only its own events/ → no write races.


2. How waiting works (polling)

There is no server, so nobody can push a notification. "Waiting" is just calling pull_global in a loop and branching on the Status. NO_NEW_PASS means sleep and try again; DOWNLOADED/ALREADY_SYNCED means go; RUN_COMPLETE means stop:

def wait_for_global(client, pass_id):       # returns DOWNLOADED | ALREADY_SYNCED | RUN_COMPLETE
    while True:
        s = client.pull_global(pass_id)     # one cheap check (+ download if new)
        if s in (DOWNLOADED, ALREADY_SYNCED, RUN_COMPLETE):
            return s                         # got it / run over -> caller proceeds or stops
        sleep(client.poll_seconds)          # NO_NEW_PASS -> e.g. 30s, then loop again

poll_seconds keeps it from hammering Azure (a tight loop would mean cost + throttling; 30s is plenty since passes take minutes). Note there is no give-up deadline under the block policy: a site waits indefinitely for the next global, and the lead waits indefinitely for every site. Liveness is handled separately by heartbeats (below) — a stalled actor is surfaced via status, not silently abandoned.

Retry is separate from polling

Polling handles "not there yet" (NO_NEW_PASS). Retry handles "the transfer itself failed" (network blip, throttling). Every pull_*/push_* wraps the actual byte copy:

def _transfer_with_retry(do_copy):          # used inside every pull_/push_ function
    for attempt in range(client.max_retries):     # default 5, from transfer.max_retries
        try:
            do_copy()                        # copy bytes + sha256 verify
            return OK
        except TransferError:
            sleep(client.retry_backoff_seconds)
    return Status.FAILED                     # exhausted retries -> caller aborts the pass

★ Insight ───────────────────────────────────── Write the marker last, so the reader can't see a half-written pass. STITCH uploads model.<ext> first, then metrics.json. pull_global treats the pass as "ready" only when metrics.json is present — so a partially-uploaded model is invisible. The marker also carries model_sha256, which the downloader re-checks against the bytes before returning DOWNLOADED. ─────────────────────────────────────────────────

The lead's wait is the same idea, gated on all expected sites (block policy) — poll the completed events until every site in directive.expected has reported, then pull_sites:

def wait_for_all_expected(client, pass_id):  # LEAD (require: all_expected)
    expected = client.read_orchestration()["directive"]["expected"]
    while not client.all_completed(pass_id, expected):     # every sites/<id>/events/pass_NNN.completed.json
        client.refresh_site_status(pass_id)  # roll events -> orchestration.site_status (heartbeats too)
        sleep(client.poll_seconds)           # a missing site shows as 'stalled' in `status`, run waits
    return client.pull_sites(pass_id)        # download all sites' pass into the local mirror

Heartbeats — telling a dead site from a slow one

While a train/aggregate hook runs, the runner refreshes heartbeat.json every heartbeat_minutes in the background. The lead folds heartbeat ages into orchestration.json's site_status, so clif-stitch status shows each site as training (fresh heartbeat), completed, or stalled (stale heartbeat) — the signal that tells an operator which site to go revive while the run blocks.

Resumable: the local mirror makes restarts free

Because the local mirror persists, re-running clif-stitch run after a crash or closed terminal just re-syncs: anything already uploaded returns ALREADY_SYNCED (skipped), and pull_global skips passes already present — so an actor resumes from its last completed pass instead of recomputing (see §3/§4).

Why polling and not Azure Event Grid / blob change events? Push notifications exist, but they need extra Azure resources and permissions — the very IT-approval cost this whole design avoids. Simple polling keeps a site to outbound HTTPS only.


3. Local site workflow (roles/local.py) — a reconciler

A local site reads the instruction set, figures out its resume point, and does only the passes it hasn't completed. It never touches global/ or orchestration.json. Re-running after a crash/closed terminal just re-enters this loop and continues.

def run_local(client, train_fn):
    k = client.resume_point()                   # 1 + highest pass I already completed (mirror+events)
    while True:
        # 1. Sync the latest global into the local mirror, OR stop if the run ended
        s = wait_for_global(client, k - 1)      # DOWNLOADED | ALREADY_SYNCED | RUN_COMPLETE
        if s is RUN_COMPLETE:                   # lead wrote global/DONE.json
            break

        # 2. Train pass k — UNLESS a crash already left a finished model on disk
        if client.has_local_model(k):           # resume: trained before dying -> just re-push
            push_local_model(client, k)          # metrics reloaded from metrics.json on disk
        else:
            client.emit(k, "started", from_global=client.global_model_path(k - 1))
            with client.heartbeating(k):        # background heartbeat.json every heartbeat_minutes
                # train_fn loads global_path, writes out_path, and RETURNS its metrics dict
                metrics = train_fn(global_path=client.global_model_path(k - 1),
                                   out_path=client.model_path(k), data_dir=client.data_dir)
            if push_local_model(client, k, metrics) is FAILED:
                raise RuntimeError(f"pass {k} upload failed")  # re-run resumes via ALREADY_SYNCED

        # 3. Report back (the "completed" event the lead waits on) and advance
        client.emit(k, "completed", **client.read_local_metrics(k))
        k += 1
    log("local site done — run complete")

train_fn is the part the site owns: it reads the global model from global_path, fits/continues it, writes the result to out_path, and returns metrics with n_samples. STITCH never sees raw data and never opens either model file. The two resume hooks — resume_point() and has_local_model(k) — are what make "pick up where it left off" automatic: on re-entry the site fetches the current pass's global (k-1), and if it already trained pass k before dying it just re-pushes rather than retraining.


4. Lead workflow (roles/lead.py) — a reconciler

The lead does everything a local site does plus init and aggregation, and is the only writer of global/ and orchestration.json. STITCH drives the pass skeleton (init, wait-for-all, sync, advance, termination) and calls the lead's own train_fn / aggregate_fn — there is no separate init hook (pass 0 is just train_fn with global_path=None), and STITCH has no FedAvg of its own. It too resumes: a published global/pass_k is skipped.

def run_lead(client, train_fn, aggregate_fn):
    N = client.num_passes

    # ---- pass 0: initialize (skip if already published) ----
    if not client.has_global_model(0):
        # there is no separate init hook: pass 0 is just train_fn with no global yet
        m0 = train_fn(global_path=None, out_path=client.global_model_path(0),
                      data_dir=client.data_dir)               # USER writes the starting model
        push_global_model(client, 0, m0)                      # upload pass 0 + orchestration -> 1

    # ---- passes 1..N ----
    for k in range(1, N + 1):
        if client.has_global_model(k):          # resume: this pass already aggregated -> skip
            continue

        # (a) lead is also a participant: train on the previous global (skip if done)
        if client.has_local_model(k):
            push_local_model(client, k)                       # re-push from disk on resume
        else:
            with client.heartbeating(k):
                metrics = train_fn(global_path=client.global_model_path(k - 1),
                                   out_path=client.model_path(k), data_dir=client.data_dir)
            push_local_model(client, k, metrics)
        client.emit(k, "completed", **client.read_local_metrics(k))

        # (b) BLOCK until ALL expected sites reported, sync them, hand paths to the aggregator
        wait_for_all_expected(client, k)        # waits indefinitely; stalled sites show in `status`
        agg = aggregate_fn(site_dir=client.sites_pass_dir(k),   # USER reads all sites/*/pass_k/ files
                           out_path=client.global_model_path(k))  # USER writes aggregated model + RETURNS metrics

        # (c) publish new global + advance the pass for everyone
        push_global_model(client, k, agg)        # also advances orchestration.json

        if agg.get("should_stop"):              # USER may signal convergence via the returned dict
            break

    # ---- terminate: ONE end marker stops every site ----
    client.write_end_marker(final_pass=k, reason="early_stopped" if k < N else "completed")
    # write_end_marker writes global/DONE.json AND sets orchestration.phase = "complete"
    log(f"federation complete after {k} passes")

aggregate_fn is the lead's code — STITCH just syncs every sites/*/pass_k/ into the local mirror and points the script at the folder; the script reads the model files itself (each paired with its metrics.json, where n_samples is the weight) and writes the combined model to out_path. FedAvg, FedProx, median — entirely the user's choice.

The lead writes exactly one global/DONE.json. Whether the run ended at N passes or on convergence, every site's next pull_global returns RUN_COMPLETE and exits — termination is centrally controlled, not counted at each site.

Resume — the same loop, re-entered

Neither loop has special resume code beyond resume_point() and the has_*_model() guards: re-running clif-stitch run simply re-enters and the idempotent checks skip finished work.

Where it died On re-run, the loop…
local, mid-train (no model_path(k)) re-pulls global k-1, retrains pass k
local, after train, before push sees model_path(k) → skips train, just pushes + reports
local, after push resume_point() is k+1 → waits for the next global
lead, after publishing global/pass_k has_global_model(k) → skips to k+1
any actor, terminal simply closed identical to a crash — re-run continues

A down site doesn't corrupt anything: the lead blocks at that pass (run waits), and the site's stalled heartbeat in clif-stitch status tells the operator to re-run it. The lead itself is revived the same way; while it's down, other sites complete their pass and wait.


5. End-to-end flow (one pass)

  LEAD                          BLOB (rendezvous)                 LOCAL SITE (e.g. NU)
  ────                          ─────────────────                 ────────────────────
  push_global_model(k-1) ─────► global/pass_{k-1}/{model,metrics}
                                orchestration.json: current_pass=k
                                       │
                                       ├──────────────────► pull_global(k-1) -> DOWNLOADED
                                       │                     emit pass_k.started + heartbeats
                                       │                     USER trains, writes model_path(k) [local data]
  (lead trains too)                    │ ◄───────────────── push_local_model(k) + emit pass_k.completed
  push_local_model(k) ───────────────► sites/NU/pass_k/{model, metrics.json}
                                sites/NU/events/pass_k.completed.json
                                       │
  wait_for_all_expected(k) ◄── ALL sites/*/events/pass_k.completed.json present?
  pull_sites(k) ◄───────── downloads sites/*/pass_k/ (verify sha256)
  USER aggregate() reads files -> writes global_model_path(k)
  push_global_model(k) ──────► global/pass_k/{model, metrics.json}
                                orchestration.json: current_pass=k+1
                                       │
                                       └──────────────────► pull_global(k)  [next pass]

No arrow goes site→site. Every arrow is a file sync (local↔blob). The lead's wait is pure reading of completed events — sites only write their own events/, so there is no conflict.


6. Why this maps cleanly to roles

Function Local site Lead
pull_global ✅ ✅
push_local_model + emit ✅ ✅ (it trains too)
pull_sites — ✅
push_global_model — ✅
writes orchestration.json / DONE.json — ✅ (single writer)
writes own events/ ✅ ✅

A local site's entire job is pull → train → push → report, in a reconciler loop branching on Status. The lead adds aggregation. That symmetry — lead = local + (wait-all, aggregate, publish, advance) — is why both roles share one small codebase, none of it ever opens a model file, and re-running either one simply resumes.


7. Local state & logs — where the checkpoint lives, and the crash trail

Everything that survives a crash is on the local disk, split across two trees that answer two different questions. When a run dies you never have to guess where you are — you read these:

<your uv project>/
├── stitch_store/<run_id>/            # CANONICAL STATE (the mirror) — THIS is the checkpoint
│   ├── global/
│   │   ├── orchestration.json        #   desired state: current_pass, phase, expected, site_status
│   │   ├── pass_NNN/{model,metrics.json}    #   globals pulled so far
│   │   └── DONE.json                 #   present IFF the whole run finished
│   └── sites/<id>/
│       ├── pass_NNN/{model,metrics.json}    #   YOUR completed outputs = the durable checkpoints
│       └── events/                   #   the structured progress trail (synced; lead reads it)
│           ├── pass_NNN.started.json
│           ├── pass_NNN.completed.json       #   sha256 + n_samples + metrics ("report back")
│           ├── pass_NNN.failed.json          #   ONE-LINE error — the marker the LEAD sees
│           └── heartbeat.json
│
└── runs/<run_id>/                    # LOCAL DEBUG TRAIL — verbose, never synced, gitignored
    ├── run.log                       #   one timestamped line per STITCH step, appended every `run`
    │                                 #     (pull / already_synced / train / push / wait / aggregate)
    ├── pass_NNN/
    │   ├── train.log                 #   FULL stdout+stderr of YOUR train()  (STITCH tees the hook)
    │   └── aggregate.log             #   lead only: FULL stdout+stderr of YOUR aggregate()
    └── last_error.json               #   machine-readable crash context, written whenever a hook raises
                                      #     {pass, phase, hook, error, traceback, model_path, log}

Two trees, because each answers a different question:

stitch_store/.../events/ (+ orchestration.json) runs/<run_id>/
Question it answers Where am I? (which pass completed) Why did it break? (the traceback)
Format small, fixed-schema JSON verbose, free-form text
Crosses the wire? yes — synced to the blob; the lead reads it no — local only, like data/
Written by STITCH (emit) STITCH tees your hook's stdout/stderr
Verbosity one line (failed.json carries error=str(e)) full traceback + everything your code printed

★ Insight ───────────────────────────────────── The split falls out of the security model, it isn't arbitrary. events/ is small, structured, and synced because the lead must read it across the blob to know a site finished — so it has to be race-free (append-only, single-writer) and can't be noisy. runs/ is verbose and local-only because it captures whatever your train() printed, which may echo cohort detail — so, like data/, it must never leave the box. It's gitignored, and it lives outside stitch_store/ (the mirror), so it isn't even a sync candidate — the whitelist never gets a chance to matter. failed.json is the one-line "it broke" the federation needs; last_error.json + train.log are the full story the operator needs. ─────────────────────────────────────────────────

Reading the checkpoint without any tooling

The resume point is reconstructable by eye: the highest sites/<id>/events/pass_NNN.completed.json is the last pass you finished, so resume_point() is that + 1. The paired sites/<id>/pass_NNN/model.<ext> is the durable model checkpoint STITCH will re-push instead of retraining (the row-2 case in §4). global/orchestration.json tells you what the lead expects next. clif-stitch status is just a human view over exactly these files.

When something happens, look here

Symptom Look at What it tells you
run just died runs/<run_id>/last_error.json which pass, which hook, full traceback, the model path it was writing
"did my pass finish?" stitch_store/.../events/pass_NNN.completed.json present ⇒ done; its sha256 / n_samples are what the lead got
"where will a re-run pick up?" highest *.completed.json + 1, or clif-stitch status the next pass to run — re-run resumes there, never pass 0
"my train() is misbehaving" runs/<run_id>/pass_NNN/train.log everything your code printed + its stderr
"the run is stuck" orchestration.json site_status (or status) which site is stalled (stale heartbeat) — go re-run that one
"is the whole run over?" stitch_store/.../global/DONE.json present ⇒ finished; reason says completed vs early-stopped

Recovery is always the same verb: fix the cause the log points to, then uv run clif-stitch run again — it reads stitch_store/ and continues from the last completed pass. Nothing in runs/ is read on resume; it is purely for humans, so deleting runs/ to clear the narrative is safe — resume still works entirely from stitch_store/.