Skip to content

Latest commit

 

History

131 Commits

Folders and files

Repository files navigation

Koochak

A tiny, hackable, function‑first training loop for PyTorch. Built to be easy to read, fork, and extend. It favors explicit functions and small modules over opaque classes or global state.

Related Projects

Koochak prepares reproducible workloads; it deliberately does not own cluster connectivity or GPU scheduling. It integrates with two small, independent projects:

  • Scruffy is an asynchronous resource scheduler for jobs running inside an existing multi-node allocation. It owns GPU reservations, dependencies, fair queueing, lifecycle state, and the event stream used by agents and dashboards.
  • Pazuzu is a resilient, site-neutral OpenSSH gateway with a typed Slurm client. It owns remote transport and is the backend for standalone Slurm jobs when no suitable Scruffy allocation is available.

The boundary is intentional: Koochak defines exactly what runs, Scruffy decides where and when it runs inside an allocation, and Pazuzu carries commands safely to a remote cluster.

Architecture at a Glance

From Experiment to Execution

flowchart TB
    config["Training config
model · data · optimization"]
    profile["Environment profile
Python · packages · compilers"]
    script["Submission script
resources · run identity"]

    prepare(["prepare_run()"])
    bundle[["PreparedRun
resolved config + immutable manifest"]]
    route{"Choose backend"}

    scruffy["Scruffy
queue inside an active allocation"]
    pazuzu["Pazuzu
standalone Slurm transport"]

    runner["Koochak runner
verify digests · preflight · exec"]
    loop(["training_loop()
iterables · DDP · AMP · hooks"])
    logs[("train.out_dir
metrics · logs · events")]
    checkpoints[("train.checkpoint_dir
numbered checkpoints · ready manifests")]
    stores["Stores file
private scheme:// → root · publish mode"]

    config --> prepare
    profile --> prepare
    script --> prepare
    prepare --> bundle --> route
    route -->|active allocation| scruffy
    route -->|standalone job| pazuzu
    scruffy --> runner
    pazuzu --> runner
    runner --> loop
    loop --> logs
    loop --> checkpoints
    stores -.->|resolves scheme://| checkpoints

    classDef input fill:#F4F0FF,stroke:#6D5BD0,color:#241C3A,stroke-width:1.5px;
    classDef core fill:#5B4B8A,stroke:#C7B9FF,color:#FFFFFF,stroke-width:1.5px;
    classDef choice fill:#FFF4E8,stroke:#D97745,color:#3A2117,stroke-width:1.5px;
    classDef scruffyNode fill:#DDF7F3,stroke:#168B83,color:#123B38,stroke-width:1.5px;
    classDef pazuzuNode fill:#FCE8DE,stroke:#D9674B,color:#44231A,stroke-width:1.5px;
    classDef runtime fill:#E8F0FF,stroke:#4977B8,color:#172D4D,stroke-width:1.5px;
    classDef result fill:#F4EDC9,stroke:#9C7B21,color:#352B10,stroke-width:1.5px;

    class config,profile,script,stores input;
    class prepare,bundle core;
    class route choice;
    class scruffy scruffyNode;
    class pazuzu pazuzuNode;
    class runner,loop runtime;
    class logs,checkpoints result;
    linkStyle default stroke:#88859A,stroke-width:1.5px;
Loading

The backend choice does not alter the prepared workload. Both paths invoke the same runner, which rejects configuration or environment drift before importing training code. Once admitted, training_loop() remains the small functional core; checkpointing, logging, distributed execution, and workload events attach through focused helpers and hooks. Logs stay in out_dir, while checkpoints go to checkpoint_dir, which can be a write-once object store named in a private stores file.

Inside training_loop()

flowchart TB
    user["User computation
model · step_fn · optimizer"]
    data["Data and policy
iterables · train_cfg · scheduler"]
    extensions["Optional extensions
hooks · eval_fn · checkpoint_dict"]

    setup["Setup once
device → shard → start hook → compile → DDP → resume from the checkpoint store"]
    batch["Next batch
optional CUDA prefetch + prepare_batch_fn"]
    micro["Micro-step × grad_accum
autocast → step_fn → scaled backward"]
    gradients["Gradient gate
unscale → finite check → clip"]
    update["Parameter update
optimizer → scaler → EMA → scheduler"]

    hooks["Observe
metrics · rank timing · hooks"]
    health["Protect
GPU health watchdog"]
    periodic["Persist
evaluation · serialize numbered checkpoint"]
    background["checkpoint_async
hand off · one publication in flight"]
    publication["Publish through the checkpoint store
create-only .pt → .ready.json → prune manifest-first"]
    announce["After commit
on_checkpoint · typed artifact event"]
    done{"max_steps reached
or data exhausted?"}
    result[["Resume-ready checkpoint
model · optimizer · RNG · EMA · config"]]

    user --> setup
    data --> setup
    extensions --> setup
    setup --> batch --> micro --> gradients --> update
    update --> hooks --> health --> periodic
    periodic -->|synchronous| publication --> announce --> done
    periodic -->|checkpoint_async| background --> done
    background -.->|background thread| publication
    done -->|next step| batch
    done -->|finished| result

    classDef input fill:#F4F0FF,stroke:#6D5BD0,color:#241C3A,stroke-width:1.5px;
    classDef setupNode fill:#5B4B8A,stroke:#C7B9FF,color:#FFFFFF,stroke-width:1.5px;
    classDef dataNode fill:#DDF7F3,stroke:#168B83,color:#123B38,stroke-width:1.5px;
    classDef compute fill:#E8F0FF,stroke:#4977B8,color:#172D4D,stroke-width:1.5px;
    classDef updateNode fill:#F4EDC9,stroke:#9C7B21,color:#352B10,stroke-width:1.5px;
    classDef boundary fill:#FCE8DE,stroke:#D9674B,color:#44231A,stroke-width:1.5px;
    classDef choice fill:#FFF4E8,stroke:#D97745,color:#3A2117,stroke-width:1.5px;
    classDef resultNode fill:#5B4B8A,stroke:#C7B9FF,color:#FFFFFF,stroke-width:1.5px;

    class user,data,extensions input;
    class setup setupNode;
    class batch dataNode;
    class micro,gradients compute;
    class update updateNode;
    class hooks,health,periodic,background,publication,announce boundary;
    class done choice;
    class result resultNode;
    linkStyle default stroke:#88859A,stroke-width:1.5px;
Loading

The user-defined step_fn owns model semantics and returns a loss plus optional metrics. Koochak owns the repetitive mechanics around it. Optional behavior is composed through functions and event hooks rather than subclasses, while the returned checkpoint captures everything required to resume the optimization state deterministically. With train.checkpoint_async, the loop keeps training while a background thread publishes the checkpoint; on_checkpoint fires once its manifest has committed, and the terminal save waits for any pending one.

Goals

  • Functional core: a single training_loop(...) with a clear, compact signature.
  • Hackable: pure-PyTorch, minimal magic, everything explicit via a small config mapping.
  • Iterable-first: data is any iterable (finite or infinite). No hidden epoch semantics.
  • Modern essentials: AMP, grad accumulation, grad clipping, logging, checkpointing, eval hooks, and DDP, each in small swappable modules.
  • Low dependency: standard library + PyTorch (OmegaConf for configs; optional: torchvision for examples, tqdm for niceties, wandb for logging).

Repository Layout

  • koochak/
    • loop.py – the core training_loop implementation (imports tiny helpers; loop remains minimal).
    • config.py – OmegaConf + dataclass config loader, defaults, and summary helpers.
    • core/
      • hooks.py – tiny hook system: merge/add/emit and rank0_only wrapper.
      • precision.py – autocast_context(mode, device) and Scaler(mode).
      • dist.py – DDP helpers: init_process_group, barrier, rank/world_size, rank0.
    • data/
      • iterable.py – to_device(batch, device), cycle(iterable), and take(iterable, n).
      • sharding.py – shard_dataset(..., mode=...), shard_iterable_dataset, shard_map_dataset.
      • shards.py – immutable dataset shards: ShardWriter, strict shard indexes, tar (WebDataset-layout) format, and plan_shards/assign_shards for per-worker reading.
      • packed.py – PackedGroups/GroupStream: each data-loading worker streams whole packs of a grouped collection (prefetch, SHA256 checks, windowed shuffle, endless passes, resume) and receives its groups' files by path; PackCache keeps cycled packs in memory; PackedGroups.read_files reads scattered single files in parallel.
    • logging/
      • stdout.py – compact TSV stdout logger + make_stdout_hooks().
      • csv.py – CSVLogger and make_csv_hooks(path).
      • jsonl.py – JSONLLogger and make_jsonl_hooks(path).
      • events.py – bounded lifecycle/progress hooks and an optional lazy Scruffy adapter.
      • wandb_logger.py – optional W&B hooks (lazy import), artifact upload.
    • jobs/
      • profile.py – strict, reusable execution-environment profiles.
      • manifest.py – deterministic config materialization and immutable launch manifests.
      • runner.py – clean-environment checks, managed child process groups, and typed output publication.
      • workflow.py – immutable PreparedTask/PreparedWorkflow models targeting Scruffy's strict workflow protocol.
      • backends.py – thin Python adapters for Pazuzu and Scruffy, including all-or-none workflow staging.
    • optim/
      • build.py – tiny builders for optimizers/schedulers (supports cosine, step, plateau, cosine_warmup).
    • storage/
      • artifact.py – immutable file/directory ready manifests with deterministic ordering, size, SHA256, provenance, and counts.
      • checkpoint.py – checkpoint publication through a Store (manifest last, optional background thread), load, auto-resume, latest, best.
      • atomic.py – atomic file writer.
      • fs.py – small FS utilities (mkdir_p, latest, best).
      • pruning.py – prune_keep_last_k(dir, pattern, k).
      • store.py – the Store protocol (write-once objects), StoreProfile, LocalStore, and open_store/register_store for pluggable backends.
      • transfer.py – copy_objects: parallel, range-splitting, verified, resumable copies between stores.
      • stores_file.py – named scheme:// stores from a private YAML file ($KOOCHAK_STORES).
      • collection.py – manifested collections: manifest.json + files.jsonl.gz, tar packs, standalone objects.
      • archive.py – archive (grouped packing, resumable), pull (subsets, verified, metadata restored), verify.
    • data/__main__.py – python -m koochak.data archive|pull|verify|ls.
      • probe.py – python -m koochak.storage.probe <location>: storage semantics and throughput report.
    • utils/
      • config.py – thin compatibility wrappers around koochak.config (get/as_dict).
      • device.py – get_device(cfg) and get_lr(optimizer).
      • seed.py – set_all_seeds(seed), make_worker_init_fn(seed), get_rng_state().
      • stats.py – SmoothedMeter, Throughput, EMA.
      • timeit.py – Timer and time_block(...) context utilities.
  • examples/mnist/
    • config.yaml – YAML-driven config split into train, data, optim, logging, wandb sections.
    • main.py – minimal end-to-end example (YAML-only CLI), stdout hooks by default, optional W&B.
  • examples/config_template.yaml – canonical template with all supported keys.
  • AGENTS.md – running design notes + TODOs for contributors.

Installation

  • Python 3.10+
  • PyTorch (CUDA optional): https://pytorch.org
  • Required: pip install omegaconf
  • Optional: pip install torchvision wandb tqdm

This repo is intentionally lightweight: it is not a packaged PyPI install. Import modules via the repo root (e.g., python -m examples.mnist.main).

Quickstart (MNIST)

  1. Configure YAML (defaults provided):

examples/mnist/config.yaml

  • train: loop behavior (max_steps, log/eval/ckpt cadence, grad_accum, amp, seed, device, out_dir, keep_last_k, ddp flag).
  • data: data_dir, batch_size, num_workers.
  • optim: optimizer + scheduler (e.g., AdamW + cosine_warmup).
  • logging: csv_path, jsonl_path.
  • wandb: set enabled: true to turn on W&B logging.
  1. Run:

python -m examples.mnist.main --config examples/mnist/config.yaml

This will download MNIST (via torchvision), print TSV logs to stdout, periodically evaluate, and write atomic checkpoints to train.out_dir (e.g., ./runs/mnist/step000000000.pt and latest.pt).

Configuration

Configs are OmegaConf-first with structured dataclass defaults. Defaults fill missing keys, user YAML overrides defaults, and CLI overrides (if any) apply last. Use OmegaConf interpolation for cross-section reuse (e.g., logging.csv_path: ${train.out_dir}/log.csv).

By default, Koochak enforces strict configuration to minimize surprises.

  • Strict mode: unknown YAML keys cause an immediate error before training. To relax, set train.strict_config: false.
  • Warnings: if strict is disabled, unknown keys print rank-0 warnings when train.config_warn_unknown: true (default true).

Example YAML toggles:

train:
  strict_config: true       # default
  config_warn_unknown: true # default, applies when strict_config=false

At startup, a brief config summary prints sections present, unknown keys (if any), and strict status.

Canonical template: examples/config_template.yaml.

In code, use koochak.config.load_config(path) and koochak.config.get_section(cfg, "train") (or similar) to access sections.

DDP sharding is explicit and opt-in:

  • train.shard_dataset: true with train.shard_dataset_mode: iterable|map to shard the training dataset.
  • train.shard_eval_dataset: true with train.shard_eval_dataset_mode: iterable|map to shard the eval dataset.
  • train.warn_unsharded: false to disable rank-0 warnings when DDP runs without Koochak sharding.

You can also shard manually in custom code via koochak.data.sharding.shard_dataset(...).

Config Map

  • train – consumed by koochak/loop.py and utilities (device, DDP/sharding, logging cadence, checkpoints, EMA, AMP).
  • data – consumed by examples or your dataset builders; use OmegaConf interpolation for shared values.
  • optim – consumed by koochak/optim/build.py for optimizer + scheduler construction.
  • logging – consumed by CLI/examples to configure stdout/CSV/JSONL hooks.
  • wandb – consumed by koochak/logging/wandb_logger.py.
  • entry – consumed by koochak/cli/train.py to import user callables.

Generic CLI

Koochak ships a generic YAML-driven CLI so you can run training without custom scripts.

Run:

python -m koochak.cli.train --config path/to/your_config.yaml

Your YAML must include an entry section to locate user code, plus the standard sections:

entry:
  model: your_pkg.model_defs:make_model     # returns nn.Module
  dataset: your_pkg.data:train_dataset      # returns iterable (or DataLoader)
  step: your_pkg.train:step_fn              # def step_fn(model, batch, ctx) -> dict
  eval_dataset: your_pkg.data:val_dataset   # optional
  eval_fn: your_pkg.train:eval_fn           # optional

train: { ... }
data:  { ... }
optim: { optimizer: {...}, scheduler: {...} }
logging: { csv_path: ..., jsonl_path: ... }
wandb: { enabled: false, project: ... }

The CLI loads config via koochak.config.load_config, prints the summary, builds the optimizer/scheduler from optim, attaches stdout/CSV/JSONL/W&B hooks, resumes from the highest valid published numbered checkpoint in train.checkpoint_dir (default train.out_dir; or starts cleanly), and calls training_loop with train_cfg. Pass --resume none to disable this lookup.

Reproducible Job Submission

Koochak compiles a training config and an execution-environment profile into an immutable launch manifest. Scruffy can enqueue that runner inside an active allocation, while Pazuzu can stage and submit it as a standalone Slurm job. Both use the same prepared run, so changing the backend does not change its configuration or environment contract. Koochak contains no hostnames, users, filesystem layout, scheduler defaults, or other site policy.

Environment profiles are strict YAML. They define an absolute Python, deterministic PATH, explicit compiler paths, required files and package versions, named secret inputs, cache directories, and built-in checks. Unknown fields and OmegaConf environment interpolation are rejected. Secret values are read only inside the worker and are never written to the profile or manifest.

version: 1
id: project-gpu-v1
python: /shared/envs/project/bin/python
environment:
  set:
    PATH: /shared/toolchains/bin:/shared/envs/project/bin:/usr/bin:/bin
    CC: /shared/toolchains/bin/gcc
    CXX: /shared/toolchains/bin/g++
    TRITON_CACHE_DIR: "{run_dir}/triton-cache"
  secrets: [TRACKING_TOKEN]
  create_directories: ["{run_dir}/triton-cache"]
requirements:
  executables: [/shared/toolchains/bin/gcc, /shared/toolchains/bin/g++]
  files: []
  packages: {torch: "2.8.0", triton: "3.4.0"}
preflight: [c_compiler, cuda, torch_compile]

For a multi-task DAG, submit_scruffy_workflow stages every PreparedRun before making one strict scruffy.submit_workflow(root, request_id, workflow_id, project_id, tasks) call. It does not signature-probe or drop unsupported artifact gates or recovery fields. During joint development, tests validate emitted specs against the exact local Scruffy checkout. Final-pin gate: before either repository is released or deployed, update Koochak's optional scruffy extra to the final reviewed Scruffy commit and rerun both full suites. Artifact-only release is represented by an empty needs list and a wait_for artifact condition. Repeated request IDs are passed unchanged so Scruffy can apply its idempotent submission semantics. Recovery mappings use Scruffy's exact v1 wire contract: max_attempts is an integer from 1 through 10, retry_on is a duplicate-free array (which may be empty), and evacuation contains exactly signal: USR1 plus a non-negative integer grace_seconds.

DeclaredOutput describes a file or directory that a consumer may wait for. The workload must create the output and atomically publish its <output>.ready.json manifest. After a successful child exit, Koochak validates all declared outputs before publishing typed workload.artifact events. A missing, corrupt, symlinked, or unpublishable output fails the runner. A SIGUSR1 request is forwarded only to that child process group and never to the parent Slurm allocation. If the leader exits while descendants remain, Koochak terminates only that owned group and fails closed before output validation. Every output is revalidated immediately before its event. Event IDs include the job, artifact ID, and content digest, so retrying after a partial spool is idempotent: an already accepted first event deduplicates while later outputs continue.

Ready-manifest and staging publication use no-overwrite hard links, stable O_NOFOLLOW descriptor reads, inode checks, and read-only regular targets. These checks defend against accidental concurrent writers and ordinary replacement races. They assume cooperative same-user shared storage: a process with permission to mutate files and directories continuously can always race a later consumer, so hostile same-UID writers require filesystem isolation or separate credentials.

Prepare a run in a committed Python submission script:

from koochak.jobs import ConfigPatch, load_environment_profile, prepare_run

prepared = prepare_run(
    name="smoke-len128",
    profile=load_environment_profile("environments/gpu.yaml"),
    python_args=["-m", "my_pkg.train", "--config", "{config}"],
    cwd="/shared/project/repo",
    run_dir="/shared/project/runs/smoke-len128",
    base_config="configs/train.yaml",
    patches=[ConfigPatch("data.max_length", 128)],
)

For a standalone Slurm job, pass backend-native resources to Pazuzu:

from pazuzu import PazuzuClient, SlurmResources
from koochak.jobs import submit_pazuzu

handle = await submit_pazuzu(
    PazuzuClient(),
    prepared,
    resources=SlurmResources(
        nodes=1,
        gpus_per_node=1,
        cpus_per_task=14,
        memory_gb_per_node=128,
        time_limit="02:00:00",
    ),
    log_dir=f"{prepared.run_dir}/logs",
)

Inside a filesystem that can see a Scruffy queue and the run directory:

from scruffy import ResourceRequest
from koochak.jobs import submit_scruffy

job = submit_scruffy(
    prepared,
    root="/shared/queues/allocation",
    resources=ResourceRequest(
        nodes=1,
        gpus_per_node=1,
        cpus_per_node=14,
        memory_gb_per_node=128,
        time_limit_seconds=7200,
    ),
    request_id="campaign/smoke/attempt-1",
    project_id="project",
)

An intermediate consumer can wait for a numbered Koochak checkpoint without reserving resources. Submit the producer in the same workflow with task_id="train":

job = submit_scruffy(
    prepared_consumer,
    root="/shared/queues/allocation",
    resources=resources,
    request_id="campaign/infer/attempt-1",
    project_id="project",
    workflow_id="campaign-1",
    task_id="infer",
    wait_for=[{
        "kind": "artifact",
        "task_id": "train",
        "artifact_id": "checkpoint/step000100000.pt",
    }],
)

The runner starts under python -I, reconstructs the environment from a small allowlist, restores scheduler-owned SLURM_*, SCRUFFY_*, and GPU identity, verifies manifest/config digests, performs the declared checks, writes preflight.json, and only then uses execve to start user code. A missing C compiler or failed Torch/Triton compilation therefore fails before model code runs. Ordinary variables, including NCCL_*, must be explicit under set; only fixed runtime identity and named secrets are inherited. See examples/jobs/ for complete site-neutral examples.

Core API

koochak/loop.py exposes:

training_loop(
  *,
  model: nn.Module,
  dataset: Iterable,                 # any iterable (finite or infinite)
  step_fn: Callable,                 # returns {"loss": Tensor, ...}
  optimizer: Optimizer,
  scheduler: Optional[_LRScheduler] = None,
  train_cfg: Mapping[str, Any],
  config_json: Optional[Mapping[str, Any]] = None,
  checkpoint_dict: Optional[Dict[str, Any]] = None,
  eval_dataset: Optional[Iterable] = None,
  eval_fn: Optional[Callable] = None,
  hooks: Optional[Dict[str, list[Callable]]] = None,
) -> Dict[str, Any]
  • step_fn(model, batch, ctx) returns {"loss": Tensor, ...}; any additional scalar values are logged.
  • ctx contains device, rank/world_size, autocast, scaler, config_json, and train_cfg.
  • The loop handles gradient accumulation, AMP, optional grad clipping, scheduler stepping (per train.scheduler_step), evaluation hooks, automatic DDP bootstrap/wrapping when train.ddp is true, and deterministic checkpointing.
  • Rank-0 prints a compact parameter count banner at startup to highlight model size changes.
  • L2 gradient clipping uses scaled reductions if the usual norm overflows. It logs grad_norm, grad_clip_coefficient, and grad_clip_scaled_norm; the ordinary path adds one scalar device-to-host synchronization. Clipping rejects NaN or infinity entries before the optimizer step. The optional nonfinite_grad_check_every check also raises rather than replacing invalid entries with zeros.
  • Atomically saves the terminal in-memory state before on_train_end and returns the same resume-ready checkpoint dictionary.

Minimal step_fn example:

def step_fn(model, batch, ctx):
    x, y = batch["x"], batch["y"]
    logits = model(x)
    loss = torch.nn.functional.cross_entropy(logits, y)
    acc = (logits.argmax(-1) == y).float().mean()
    return {"loss": loss, "acc": acc}

Hooks and Logging

  • Create hooks by event name: {"on_log": [fn], "on_eval_end": [fn]}.
  • Hook events emitted by the loop include on_train_start, on_step_end, on_log, on_eval_end, on_checkpoint, on_train_end, and on_exception.
  • Built-in hooks:
    • koochak.logging.stdout.make_stdout_hooks() – TSV prints; rank-0 only.
    • koochak.logging.csv.make_csv_hooks(path) – append metrics to CSV; rank-0 only.
    • koochak.logging.jsonl.make_jsonl_hooks(path) – one JSON per line; rank-0 only.
    • koochak.logging.events.make_event_hooks(publish) – rank-0 workload.phase, workload.progress, workload.milestone, and workload.artifact events for an external coordinator. Training progress defaults to approximately one event every 30 seconds at completed-step boundaries and includes completed/total steps; evaluations and checkpoint references are always attempted. Numbered checkpoints include a strict publication record (stable artifact ID, absolute path, size, SHA256, and ready-manifest path) which can satisfy a declared Scruffy artifact condition. Payloads never include full resolved configs or checkpoint contents.
    • koochak.logging.events.make_scruffy_hooks() – requires SCRUFFY_ROOT and SCRUFFY_JOB_ID; the exact Scruffy client is imported and validated when hooks are constructed, before training starts. Declare scruffy in the execution profile and include its shared source/site root in that profile's explicit PYTHONPATH; the isolated runner checks that import during preflight. Install koochak[scruffy] only when the compatible client is available from the target environment. Checkpoint acknowledgement waits up to 300 seconds by default. Override with KOOCHAK_SCRUFFY_ARTIFACT_ACK_TIMEOUT_SECONDS or pass artifact_ack_timeout_s=<seconds> (the explicit argument takes precedence). Only strict numbered workload.artifact checkpoint publications use wait=True; lifecycle and evacuation milestone events remain asynchronous. A rejected or conflicting strict checkpoint acknowledgement fails closed. An acknowledgement timeout is reconciled against the durable receipt and journal-derived job artifact evidence without publishing a second event. If the bounded reconciliation deadline still expires, Koochak raises the checkpoint-safe retryable reason checkpoint_ack_timeout, preserves the just-written numbered checkpoint for resume: auto, and the managed runner returns exit code 76. The Scruffy deployment must map that code to its retryable checkpoint reason and enforce the task's capped max_attempts; older Scruffy versions record it as application_exit. Failed publication never releases a dependent task; Scruffy leaves it blocked rather than inferring readiness from the filesystem.
    • koochak.logging.wandb_logger.make_wandb_hooks(cfg) – W&B logging/artifacts; rank-0 only.
  • Stdout and W&B record the resolved config at on_train_start; CSV/JSONL remain metric logs.
  • Compose hooks with koochak.core.hooks.merge(a, b). Gate any custom hook via koochak.core.hooks.rank0_only(fn) to ensure single-emission under DDP.
  • The generic python -m koochak.cli.train entrypoint automatically merges the Scruffy hooks when both worker variables are present. Publisher errors warn once and remain non-fatal; CSV, JSONL, W&B, and raw training logs remain the detailed telemetry sources.

YAML-driven logging (example):

logging:
  csv_path: ./runs/mnist/log.csv
  jsonl_path: ./runs/mnist/log.jsonl
wandb:
  enabled: false

If csv_path/jsonl_path are omitted, the MNIST example defaults to <train.out_dir>/log.csv and <train.out_dir>/log.jsonl.

W&B artifacts:

  • The W&B hook versions checkpoints as a single artifact per run named <prefix>-<run_id> (default prefix model).
  • Each upload includes aliases: latest, step-<n>, and when improved metrics are seen, best and best-<metric>.
  • Config overrides (optional) under wandb:
    • artifact_name_prefix (str, default model)
    • artifact_type (str, default model)

Restartable jobs that actually select a checkpoint with resume=auto must set wandb.resume: allow and provide a stable wandb.name or explicit wandb.id. The first attempt may create the run without these restart-only requirements; Koochak enforces them only when a later attempt finds durable checkpoint state.

EMA Support

  • Enable EMA by setting train.ema.enabled: true (or by providing decay/profile keys while enabled is unset). Nested config lives under train.ema.*; legacy flat keys (ema_decay, ema_eval, etc.) are still honored.
  • Supported options: decay, decay_init, warmup_steps, schedule (constant, linear, cosine), profile (constant, power), gamma/srel for power-law schedules, offload_to_cpu, pin_memory, update_every, compensate_update_every, and eval_with_ema to run eval with shadow weights.
  • Thinned EMA updates are decay-compensated by elapsed model steps. For example, update_every: 1 uses decay, while update_every: 2 uses decay ** 2 on each EMA update. The compensate_update_every key is retained for config/checkpoint compatibility; compensated behavior is the implementation.
  • Dual EMA tracking is available via train.ema.dual.enabled plus gamma1/gamma2 or srel1/srel2; both shadows are saved and restored from checkpoints.
  • For EDM2-style post-hoc EMA tuning, collect the two dual EMA states from multiple saved checkpoints and pass the flattened list to koochak.utils.ema_posthoc.reconstruct_power_ema_state_dict(...). reconstruct_dual_power_ema_state_dict(...) remains the lightweight same-checkpoint two-shadow helper.
  • EMA state is serialized alongside the model (and matches state-dict prefixes automatically) so resumes and manual loads stay seamless.

Checkpointing

  • Normal completion writes step{N:09d}.pt, where N is the number of completed optimizer updates. Terminal checkpoints record step=N and next_step=N; periodic checkpoints retain their zero-based update index and record next_step=step+1. This keeps filenames and resume cursors unambiguous when a run is extended.

  • Checkpoints go to train.checkpoint_dir, which defaults to train.out_dir and may be a directory or a named scheme:// URI from the stores file (see below). Logs, W&B files, and GPU-health summaries always stay in out_dir: they append, which write-once object-storage mounts refuse. So a run can keep its logs on a POSIX filesystem and its checkpoints on object storage:

    train:
      out_dir: /posix/runs/exp0              # logs
      checkpoint_dir: archive://runs/exp0    # checkpoints, via the stores file
      checkpoint_async: true
  • Every save goes through the checkpoint store (koochak.storage.checkpoint.publish(store, name, data, keep_last_k=k)): the numbered checkpoint is written create-only and synced, then <path>.ready.json commits it with its stable artifact ID, byte size, and SHA256. Re-saving an existing step first uncommits it (manifest, then file). Pruning keeps the last k pairs and deletes each manifest before its checkpoint. Directory stores also keep a latest.pt symlink; write-once stores (publish: exclusive) get none, because the fallback would be a full copy of every checkpoint. save(ckpt, path, keep_last_k) wraps publish for a directory path.

  • train.checkpoint_async: true serializes each periodic checkpoint at once and publishes it on a background thread, with at most one publication in flight, so rank 0 does not stall on slow storage. on_checkpoint fires when the manifest has committed (the hook's tensors are live by then; read the saved file for exact values), a failed background publication fails training on its next step, and terminal, evacuation, and GPU-health checkpoints drain the pending publication and then save synchronously. Code that reads a checkpoint file as soon as _save_checkpoint returns needs synchronous saves.

  • Bind Scruffy dependencies to the immutable numbered artifact ID, such as checkpoint/step000100000.pt, never to the mutable latest.pt alias.

  • koochak.storage.checkpoint.load(path) loads to CPU.

  • koochak.storage.checkpoint.latest(location) returns latest.pt if present or the most recent step checkpoint.

  • koochak.storage.checkpoint.highest_valid_published(location) selects the highest numbered checkpoint whose exact ready manifest, path, size, SHA256, and resume cursor validate. It never treats latest.pt, a directory, or incomplete/corrupt scaffolding as evidence. Reads go through the checkpoint store, so a mount with read_settle_seconds waits out files another node has just closed. A missing or invalid candidate falls back to a lower one; any other I/O error is raised, so an unreadable newest checkpoint never silently rolls training back.

  • koochak.storage.checkpoint.resolve_auto_resume(location) returns (path, checkpoint) from the same single payload read whose manifest, path, size, SHA256, and cursor were validated. Use checkpoint["next_step"] before constructing a dataset, then pass both values as checkpoint_dict and auto_resume_path to training_loop(..., resume="auto") to retain selection events and typed artifact republishing.

  • koochak.storage.checkpoint.best(location, key) selects the lowest metric across checkpoints.

  • Every location above is a directory, a scheme:// URI, or a Store (checkpoint_store(location) resolves it).

Safe evacuation and resume

Evacuation is opt-in. Set train.evacuation_enabled: true (and optionally train.evacuation_signal: USR1) or pass an EvacuationController/True through the training_loop(..., evacuation=...) API. Koochak installs a minimal SIGUSR1 handler that only records an in-memory request. At the end of a completed optimizer update, all DDP ranks reconcile that request; rank 0 then atomically publishes a numbered terminal checkpoint with next_step set to the first unexecuted update, emits the artifact and evacuation milestone events, finishes hooks, and exits with reserved code 75. The handler performs no I/O or distributed operations.

Use training_loop(..., resume="auto") or the CLI's default --resume auto to load only the highest valid published numbered checkpoint. With no valid checkpoint, auto-resume starts from step zero, making the same immutable command safe on both a first attempt and a retry. latest.pt remains a convenience pointer and is not resume evidence.

DDP compatibility:

  • The loop saves the underlying module weights when the model is wrapped in DistributedDataParallel (i.e., uses model.module.state_dict()), making checkpoints portable across single-GPU and DDP.
  • When loading manually, use the provided helpers if your loading target differs in wrapping:
    • from koochak.storage.checkpoint import match_state_dict_to_model
    • target = getattr(model, 'module', model)
    • target.load_state_dict(match_state_dict_to_model(target, ckpt['model']))

Storage Backends and Dataset Shards

Koochak is moving its persistence onto one narrow interface so the same code runs on parallel filesystems, object stores, and FUSE mounts over object storage. See specs/storage-abstraction.md for the full design and phases. Checkpoints are published through a Store (see Checkpointing above); splitting them into parallel parts and replicating them between tiers remain.

flowchart TB
    storesfile["Stores file
private YAML · $KOOCHAK_STORES"]
    plugins["Plugin stores
koochak.stores entry points"]
    probe["Storage probe
semantics · settings · profile"]
    open(["open_store(location)"])
    local["LocalStore
link: POSIX · exclusive: object-storage FUSE
fsync · read-back · read_settle_seconds"]
    store[["Store protocol
get · open · put create-only · stat · list · delete"]]

    checkpoints["Checkpoints
publish · prune · resolve_auto_resume"]
    shards["Dataset shards
ShardWriter · plan_shards · assign_shards"]
    transfer["copy_objects
parallel · ranged · SHA256 · resumable"]
    collections[("Collections
manifest.json · files.jsonl.gz · packs · objects")]
    tool["python -m koochak.data
archive · pull · verify · ls"]

    probe -.->|recommends| storesfile
    storesfile --> open
    plugins --> open
    open --> local --> store
    open -->|custom backend| store
    store --> checkpoints
    store --> shards
    store --> transfer --> collections
    tool --> transfer

    classDef input fill:#F4F0FF,stroke:#6D5BD0,color:#241C3A,stroke-width:1.5px;
    classDef core fill:#5B4B8A,stroke:#C7B9FF,color:#FFFFFF,stroke-width:1.5px;
    classDef runtime fill:#E8F0FF,stroke:#4977B8,color:#172D4D,stroke-width:1.5px;
    classDef toolNode fill:#DDF7F3,stroke:#168B83,color:#123B38,stroke-width:1.5px;
    classDef result fill:#F4EDC9,stroke:#9C7B21,color:#352B10,stroke-width:1.5px;

    class storesfile,plugins,probe input;
    class open,store core;
    class local,checkpoints,shards,transfer runtime;
    class tool toolNode;
    class collections result;
    linkStyle default stroke:#88859A,stroke-width:1.5px;
Loading

Every layer above the protocol is backend-neutral: the same checkpoint, shard, and archive code runs against a parallel filesystem or an object-storage mount, and the stores file carries the site-specific settings.

Training on grouped collections. Archive a dataset with --groups (path,group,order) so every file one training record needs is stored contiguously, then stream it:

from koochak.data.packed import GroupStream, PackCache, PackedGroups
from koochak.storage.store import open_store

store = open_store("archive://datasets/my-data/v1")
packed = PackedGroups.load(store)             # file table, groups, packs
plan = packed.plan(num_owners=world_size * workers_per_rank, seed=0)
stream = GroupStream(store, packed, plan[owner], seed=epoch_seed, window=2, prefetch=2,
                     cache=PackCache(4 << 30), retries=3)
for group in stream:                          # endless; a seeded shuffle per window of packs
    payload = group.files["some/original/path.npz"]

Each worker owns whole packs (balanced by group count), fetches them whole in a background thread, verifies them against the manifest, and yields their groups shuffled within windows of window packs. start=n skips n groups without reading the skipped packs. Plans for disjoint selections (select=) can be combined to balance several kinds of groups separately.

Reads outside the streams, such as a small sidecar per record while the dataset is built, belong in one packed.read_files(store, entries, streams=32) call: it reads only those files' bytes (neighbours share a request), never their whole groups, and opens each pack once. On an object-storage mount an open can take a second or more and the mount admits only a few per second, so opening once per pack rather than once per file is several times faster.

  • koochak.storage.store.Store holds write-once objects under relative keys: get (whole or byte range), open, put (create-only; returns once the bytes read back), stat, list(prefix), delete, and optional local_path. There is no rename, append, overwrite, or symlink in the contract, because object storage cannot provide them atomically.
  • open_store(location) maps plain paths and file:// URIs to LocalStore, and any other scheme:// to a factory registered with register_store or exposed by an installed package under the koochak.stores entry-point group. Site-specific backends live in private packages, not in this repository.
  • Stores file. Site-specific facts stay out of code: a private YAML file maps scheme names to roots, publish modes, and measured profiles, and open_store("archive://datasets/foo") returns a LocalStore rooted at <root>/datasets/foo with those settings. The same relative path under two schemes names the same data on two tiers. The file is $KOOCHAK_STORES, else ~/.config/koochak/stores.yaml (XDG); values resolve with OmegaConf (${oc.env:USER}), unknown keys fail, and it holds no secrets. See examples/storage/stores.example.yaml; keep the real file private.
  • LocalStore(root, publish="link") publishes via a hidden temp file plus a hard link. On mounts without hard links use publish="exclusive" (O_EXCL create, published on close); fsync=False and verify_readback=True, settle_seconds=... cover mounts that reject fsync or close asynchronously. read_settle_seconds retries reads (open, get, stat) that fail with ETIME or EIO, which some object-storage mounts return for minutes after another node closed a file, or with a lost connection (ECONNABORTED, ECONNRESET, ENOTCONN) while the mount's client is cut off from the service. Writes are never retried. The job runner's required-file preflight waits out the same errors for 15 minutes.
  • StoreProfile records a store's measured performance (per-request cost, cold bandwidth per stream, useful concurrency, whether byte ranges of one object scale, part size, whether listing is acceptable). LocalStore(..., profile=...) attaches one; profile_of(store) returns it (a conservative remote default otherwise), and min_object_bytes(profile) gives the size below which files should be packed rather than stored individually.
  • koochak.storage.transfer.copy_objects(source, target, items) is the one way bytes move between stores: streams items at once, large objects read as parallel byte ranges into one create-only put, SHA256 checked in flight against manifest digests (a mismatching new object is removed), existing targets of the expected size skipped so reruns resume, first failure raised. Defaults come from the stores' profiles.
  • python -m koochak.storage.probe <path-or-uri> [--json] checks those semantics (exclusive create, rename, hard links, symlinks, fsync, in-place writes, visibility of unclosed files, read-back delay), recommends LocalStore settings and a StoreProfile, and measures small-object latency, listing, cold versus cached single-stream, concurrent, and range-parallel reads (page caches are dropped before cold reads where the OS allows). --emit-profile prints just the recommended profile: block to paste into the stores file. --checkpoint-bytes 4G --checkpoint-parts 1,8 adds a checkpoint-sized write as 1..N concurrent parts with read-back timing; --dataset-shards 32 --shard-bytes 256M --readers 1,4,16 builds a synthetic dataset with ShardWriter and reads it with N spawned processes via assign_shards (cold and warm passes). Run it on a compute node; it cleans up after itself.

Collections and the data tool (koochak.storage.collection, koochak.storage.archive):

  • A collection is many files under one prefix, described by manifest.json (written last; the commit point) and files.jsonl.gz (per file: path, size, SHA256, mode, exact mtime, group, and either its pack and byte offset or its standalone object). Small files live in plain tar packs (packs/pack-NNNNNN.tar), large ones as objects/<path>. Listing and stat come from the manifest, never from listing storage.
  • python -m koochak.data archive SRC DEST [--groups groups.csv] packs a local tree. The groups table (path,group[,order]) keeps each group's files contiguous in one pack, in order, so a worker loads a group with one range read; unlisted files group by directory. Files at least --object-bytes (default from the target's profile) become objects; --layout objects stores every file as-is (e.g. already-sharded datasets). Reruns resume from per-pack records and refuse leftovers that do not match the plan. --dry-run reports files, bytes, packs, objects, and request counts.
  • python -m koochak.data pull SRC DEST [--include GLOB] restores all or some files with merged range reads, checks every SHA256, restores mode and exact mtime, and skips files already present. verify [--deep] checks sizes or every byte; ls [--include GLOB] [--long] lists from the manifest.
  • --include GLOB keeps only matching files while walking (excludes still prune directories); --skip-symlinks skips links instead of failing; --min-age-hours H skips files modified in the last H hours (e.g. checkpoints a live run may still use). Skips are counted in the report.
  • --files-from LIST archives exactly the listed paths (relative to SRC) without walking SRC. --delete-source makes the archive a move: after the committed collection passes a deep verify, sources whose size and mtime still match the manifest are deleted; failed verification deletes nothing and changed files are kept and reported. The manifest metadata records the source location.
  • Symlinks and special files are refused (exclude them with --exclude).

Datasets (koochak.data.shards):

  • Build a dataset with ShardWriter(store, "datasets/foo", target_bytes=...): records are packed into ~target_bytes shards and index.json is written last, so an interrupted build never commits. The index lists every shard's key (relative to the index), size, SHA256, and record count; its own SHA256 identifies the dataset version. TAR stores {"__key__": k, ext: bytes} samples in the WebDataset layout with deterministic headers.
  • load_index, read_shard(..., verify=True), describe_shard, and write_index read, verify, and index shards (including pre-existing files).
  • plan_shards(index, num_workers=N, epoch=e, seed=s) deals whole shards to all N data-loading workers before anything is read: each shard goes to exactly one worker per epoch, workers get contiguous runs balanced by record count, and every process computes the same plan. It raises on every worker alike if a worker would get nothing. worker_coordinates(rank=..., world_size=...) returns the global worker id inside a DataLoader worker; capture rank and world size in the main process when building the dataset.

Resuming Across DDP/Single GPU

When moving between single-GPU and DDP runs, key prefixes can differ (module.). The loop saves the underlying module weights for portability, but if you’re loading manually, use the helpers:

from koochak.storage.checkpoint import load, match_state_dict_to_model

ckpt = load(path)
target = getattr(model, 'module', model)
state = match_state_dict_to_model(target, ckpt['model'])
target.load_state_dict(state)

Checkpoint dict fields include: step, model, optimizer, scheduler (optional), scaler (optional), config, RNG state, wall_time, and metrics.

Distributed (DDP)

  • When train.ddp: true, the loop auto-initializes the process group (if needed), pins the model to the local device, and wraps it in torch.nn.parallel.DistributedDataParallel. Pass train.find_unused_parameters: true if you need the corresponding DDP flag.
  • Sharding is explicit: use train.shard_dataset/train.shard_dataset_mode (and the eval equivalents) or call koochak.data.sharding.shard_dataset(...) in custom code. If DDP is enabled and datasets are not marked as sharded, rank 0 emits a warning by default.
  • If you prefer manual control, initialize ahead of time via koochak.core.dist.init_process_group(...); the loop will detect the existing group and skip auto-init.
  • Launch with torchrun as usual:

torchrun --nproc_per_node=8 -m examples.mnist.ddp_main --config examples/mnist/config.yaml

The DDP launcher:

  • Calls init_process_group(backend=...) and sets the current CUDA device from LOCAL_RANK.
  • Forces train.ddp: true and keeps other training settings from YAML.
  • Uses the same logging configuration (stdout/CSV/JSONL and optional W&B).

Optimizers and Schedulers

  • koochak.optim.build.build_optimizer(params, cfg) supports adamw, adam, sgd.
  • koochak.optim.build.build_scheduler(optimizer, cfg, train_cfg) supports cosine, step, plateau, and cosine_warmup.

Example YAML (snippets):

optim:
  optimizer:
    name: AdamW
    lr: 0.0003
    weight_decay: 0.01
  scheduler:
    name: cosine_warmup
    warmup_steps: 100
    T_max: null   # falls back to train.max_steps
    eta_min: 0.0

Reproducibility

  • koochak.utils.seed.set_all_seeds(seed) sets Python/NumPy/Torch seeds and (optionally) CUDA seeds.
  • koochak.utils.seed.make_worker_init_fn(seed) seeds DataLoader workers deterministically.
  • RNG state is stored in checkpoints (get_rng_state()), so resumed runs continue deterministically.
  • On resume, the training loop restores RNG state from the checkpoint (Python/NumPy/Torch CPU/CUDA) before resuming steps, so randomness inside step_fn (e.g., torch.rand) is reproducible across restarts.
  • In DDP, prefer per-rank seeding (e.g., set_all_seeds(seed + rank)) and rank-aware worker seeding (make_worker_init_fn(seed, rank=rank)) to avoid correlated randomness. Checkpoints saved on rank 0 include per-rank RNG states and are used on resume to restore each rank’s RNG deterministically.

Contributing

  • Start with README.md and skim design_doc.md to understand the philosophy.
  • See AGENTS.md for current implementation notes and a living TODO list. Keep it up to date as you work.
  • Code style: clear, minimal, single-purpose modules. Favor functions and plain dicts over classes.

Tests

Unit tests live under tests/. Use your preferred runner (e.g., pytest) from the repo root:

pip install pytest
pytest -q

License

MIT (see LICENSE).

About

Intensely minimal and hackable Pytorch training utilities

Resources

Stars

4 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages