Skip to content
smturtle2Public

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Latest commit

 

History

108 Commits

Folders and files

Repository files navigation

cr-train

One-class training toolkit for satellite cloud removal -- deterministic sampling, fast HF block streaming, and DDP out of the box.

Python 3.12+ PyTorch 2.4+ Dataset MIT License

English | 한국어


Table of Contents

Highlights

  • Single-class API -- Trainer with step() + test(), nothing else to learn
  • Deterministic block sampling -- one-seed system (seed) for exact reproducibility
  • HF v2 block data -- streams compressed .crpack blocks or persists them locally
  • Distributed training -- automatic DDP wrapping, rank-aware block partitioning, all-reduce metrics
  • JSONL experiment tracking -- every train/validation epoch and startup event recorded to metrics.jsonl
  • Zero config data -- streams directly from Hermanni/sen12mscr-v2; no manual download needed

Quick Start

Installation

uv add git+https://github.com/smturtle2/cr-train.git

Minimal example

from cr_train import Trainer
import torch
from torch import nn
from torch.nn import functional as F

class FusionBaseline(nn.Module):
    def __init__(self):
        super().__init__()
        # 2 SAR channels + 13 cloudy optical channels = 15 input channels
        self.body = nn.Sequential(
            nn.Conv2d(15, 64, 3, padding=1), nn.GELU(),
            nn.Conv2d(64, 64, 3, padding=1), nn.GELU(),
            nn.Conv2d(64, 13, 1),  # 13 target optical channels
        )

    def forward(self, sar, cloudy):
        return self.body(torch.cat([sar, cloudy], dim=1))

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = FusionBaseline().to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)

trainer = Trainer(
    model, optimizer,
    loss=lambda pred, batch: F.l1_loss(pred, batch["target"]),
    metrics={"mae": lambda pred, batch: torch.mean(torch.abs(pred - batch["target"]))},
    max_train_samples=2048,
    max_val_samples=256,
    max_test_samples=256,
    batch_size=4,
    accum_steps=4,  # 4 micro-batches per optimizer update
    epochs=2,
    seed=42,
    output_dir="runs/sen12mscr",
    train_crop_size=128,
    train_random_flip=True,
    train_random_rot90=True,
)

for _ in range(trainer.epochs):
    print(trainer.step())

print(trainer.test())
Expected output
train  ░░░░░…░░█░…░░░█░…  32 blocks (2 048 rows)
val    ░░░░░░░░░░░█░░░░░…   4 blocks (  244 rows)

Epoch   1/  2 │ train │ loss    0.0423 │ mae    0.0312 │ lr      1e-3 │ 12.3s
              │ val   │ loss    0.0391 │ mae    0.0298 │              │ 2.4s

Epoch   2/  2 │ train │ loss    0.0390 │ mae    0.0294 │ lr      1e-3 │ 11.8s
              │ val   │ loss    0.0372 │ mae    0.0281 │              │ 2.3s

              │ Test  │ loss    0.0390 │ mae    0.0290 │              │ 2.1s

Examples

CLI training

Run the bundled training script with the built-in FusionBaseline model:

uv run python examples/train_sen12mscr.py \
  --max-train-samples 2048 \
  --max-val-samples 256 \
  --max-test-samples 256 \
  --batch-size 4 \
  --accum-steps 4 \
  --grad-clip-norm 1.0 \
  --epochs 2 \
  --scheduler warmup-cosine \
  --scheduler-timing after_validation \
  --warmup-epochs 1 \
  --train-crop-size 128 \
  --train-random-flip \
  --train-random-rot90 \
  --output-dir runs/sen12mscr-example

Pass --max-train-samples none (or full) to bypass sampling and train on the entire split. By default the script streams HF v2 blocks. Pass --no-streaming --dataset-dir PATH to persist selected blocks locally and reuse a complete split without revalidation on later runs. Use --accum-steps N to accumulate gradients across N micro-batches before each optimizer update; global_step in get_state() and checkpoints counts optimizer updates, not micro-batches. Train batches use a deterministic sample-mixed order across active blocks, while validation and test keep the selected block order. Pass --grad-clip-norm VALUE to clip gradients before each optimizer update. Non-finite train losses and gradients fail fast before optimizer.step(). The bundled script uses a custom WarmupCosineScheduler subclass by default to show the public scheduler API end-to-end; pass --scheduler none to disable it. Use --scheduler-timing after_validation|before_optimizer_step|after_optimizer_step to control when Trainer calls scheduler.step(). The bundled warmup-cosine example stays on the default epoch-based after_validation timing. The training augmentations apply only to the train split; validation and test stay at the original 256x256.

Sampling algorithm visualization

See how the uniform exact-k block-selection bitmask is built for partial requests. Full-split requests report planner_mode="full_split" and select every block in order:

uv run python examples/bitmask_sampling_demo.py \
  --total-rows 107072 \
  --requested-rows 2048 \
  --seed 9

Output shows the raw draw order, the final selected block indices, and a bitmap of selected (■) vs. skipped (□) logical blocks.


What Trainer Handles

Most users only need from cr_train import Trainer. Once you construct it, Trainer automatically:

  • resolves the HF v2 manifest and split catalogs
  • streams or locally prepares only the splits needed for the current call
  • builds iterable dataloaders, rank-aware block partitioning, and sample-mixed train order
  • writes metrics.jsonl
  • shows running-average loss and metrics with batch-level tqdm during training, then prints aligned train/validation summary lines with per-split elapsed time
  • appends metrics to metrics.jsonl in the output directory (one JSON object per line)

You do not need any dataset download or dataloader setup code for the normal training flow. Persistence and inference are explicit: call save_checkpoint(), load_checkpoint(), save_weights(), load_weights(), predict(), and get_state() when you need them.


API Reference

Trainer.__init__

Parameter Type Default Description
model nn.Module (required) PyTorch model. forward(sar, cloudy) signature, returns a prediction tensor.
optimizer Optimizer (required) Must be constructed from model.parameters().
loss Callable (required) (prediction, batch) -> scalar tensor.
metrics dict[str, Callable] None {"name": (prediction, batch) -> scalar}. Logged per epoch.
scheduler LRScheduler | None None Optional scheduler built from the same optimizer. Standard schedulers step once after validation. ReduceLROnPlateau is also supported.
scheduler_timing str "after_validation" When Trainer calls scheduler.step(). Supported values: after_validation, before_optimizer_step, after_optimizer_step. Keep epoch-based schedulers such as the bundled warmup-cosine example on the default after_validation timing.
scheduler_monitor str | None None Monitor path for ReduceLROnPlateau. Default is val.loss. Supported values are val.loss and val.metrics.<name>.
max_train_samples int | None None Requested train rows. Converted to 32-row HF v2 block count. None = full split.
max_val_samples int | None None Same for validation.
max_test_samples int | None None Same for test.
batch_size int 4 Batch size for all DataLoaders.
accum_steps int 1 Number of micro-batches to accumulate before each optimizer update. global_step counts these optimizer updates, not individual micro-batches.
epochs int 1 Total training epochs. Call step() once per epoch.
seed int 42 Seed controlling deterministic block selection and epoch-wise train sample order.
output_dir str | Path "runs/default" Directory for metrics.jsonl and the default save_checkpoint() / save_weights() output files.
streaming bool True Stream compressed HF v2 blocks through a shared staging pipeline. Consumed staged blocks are deleted automatically.
dataset_dir str | Path | None None Persistent local HF v2 dataset directory used when streaming=False (None = ~/.cache/cr-train/sen12mscr-v2). Complete full splits are trusted without revalidation; partial/incomplete local splits verify selected block headers and refill missing or bad blocks.
num_workers int | "auto" "auto" Job-level PyTorch DataLoader worker budget. "auto" resolves to min(16, max(1, os.cpu_count() // 3)); under DDP this budget is divided across ranks with at least one worker per rank when the budget is nonzero.
multiprocessing_context str | None None Explicit worker start method. When num_workers > 0 on CUDA, Trainer defaults this to "spawn" for safer worker startup.
train_crop_size int | None 128 Apply random square crops of this size to train batches before they leave the collate step.
train_random_flip bool True Apply independent random vertical/horizontal flips to train batches.
train_random_rot90 bool True Apply random 0/90/180/270 degree rotations to train batches.
grad_clip_norm float | None None Optional max grad norm applied before each optimizer update. Non-finite train losses and gradients always fail fast before stepping.
mixed_precision str "off" Autocast mode: off, bf16, or fp16. fp16 uses GradScaler on CUDA; bf16 uses autocast without scaling.

Trainer.step() -> dict

Runs one training epoch + validation. Returns:

{
    "epoch": 1,
    "train": {
        "loss": 0.0423,
        "metrics": {"mae": 0.0312},
        "lr": [0.001],
        "num_samples": 2048,
        "num_batches": 512,
        "samples_per_sec": 142.3,
        "batches_per_sec": 17.8,
    },
    "val": {
        "loss": 0.0391,
        "metrics": {"mae": 0.0298},
        "num_samples": 256,
        "num_batches": 64,
    },
    "elapsed_sec": 12.3,
}

Trainer.test() -> dict

Runs test evaluation with the current model state. Returns:

{
    "epoch": 2,
    "loss": 0.0387,
    "metrics": {"mae": 0.0295},
    "num_samples": 256,
    "num_batches": 64,
}

Trainer.save_checkpoint(path: str | Path | None = None) -> Path

Writes a resumable checkpoint containing model, optimizer, epoch, and global_step. global_step counts completed optimizer updates rather than individual forward/backward micro-batches. If a scheduler is configured, its state is included under scheduler. When path is omitted, the file is written to <output_dir>/epoch-XXXX.pt using the current completed epoch.

Trainer.load_checkpoint(path: str | Path) -> dict

Restores model, optimizer, epoch, and global_step from a checkpoint file and returns: For accumulation runs, the restored global_step is still the count of completed optimizer updates. If both the trainer and the checkpoint have a scheduler state, that state is restored too. Older checkpoints without scheduler remain loadable, but scheduler progression is not reconstructed automatically in that case.

{
    "path": Path("runs/sen12mscr/epoch-0005.pt"),
    "epoch": 5,
    "global_step": 2560,
}

Trainer.save_weights(path: str | Path | None = None) -> Path

Writes model weights only. When path is omitted, the file is written to <output_dir>/model-epoch-XXXX.pt.

Trainer.load_weights(path: str | Path, *, strict: bool = True) -> None

Restores model weights without touching optimizer state, scheduler state, or runtime counters. Accepts either a weights-only file from save_weights() or a checkpoint file from save_checkpoint().

Trainer.predict(batch: Mapping[str, Any]) -> Any

Runs a single forward pass in eval() + no_grad() mode and then restores the previous training mode. batch should provide at least sar and cloudy tensors.

Trainer.get_state() -> dict

Returns the current runtime state. global_step is the number of completed optimizer updates.

{
    "epoch": 5,
    "epochs": 10,
    "global_step": 2560,
    "lr": [0.0005],
    "device": torch.device("cuda:0"),
    "distributed": False,
}

Advanced: block planner inspection

To inspect the deterministic uniform exact-k block planner directly, use the low-level surface under cr_train.data:

from cr_train.data import BLOCK_SIZE, trace_plan_sample

See examples/bitmask_sampling_demo.py for a full visualization.


Architecture

Trainer defaults to streaming=True against the HF v2 block dataset. Layout-v15 stores 32-row compressed .crpack blocks and records startup events in metrics.jsonl. Partial requests use deterministic uniform exact-k logical block selection keyed by seed, while full-split requests bypass sampling and select every block in order. Training sample order still changes by epoch through seed + epoch_index, but samples are mixed across active blocks instead of draining one block at a time.

In streaming mode a producer process downloads selected HF v2 blocks into a shared staging directory under ~/.cache/cr-train/streaming-stage. Staging capacity is split into deterministic worker-owned queues, so each worker's next blocks stay eligible for download even when the staging buffer is full. Consumed staged blocks are deleted immediately, keeping network download and sample decoding pipelined without requiring a persistent local copy.

With streaming=False, dataset_dir points to a persistent local HF v2 mirror. step() prepares train, validation, and test once up front; later epochs reuse the prepared split state. Complete full splits are trusted without revalidation. If a local split is partial or incomplete, only the selected block headers are checked and missing or bad blocks are fetched from Hugging Face.

  • Startup output prints a one-line ■/□ block timeline on completion.
  • Equal seed values keep the same uniform exact-k block-selection membership for partial requests; full-split requests always include every block.
  • CUDA runs with worker processes default to the safer "spawn" multiprocessing context once num_workers > 0.
  • Persistent local datasets are never auto-deleted. Remove dataset_dir manually to reclaim disk space.

Distributed Training

Use setup_distributed_from_env() before moving the model to a device, then Trainer auto-wraps the model in DistributedDataParallel:

uv run torchrun --standalone --nproc-per-node=2 examples/train_sen12mscr.py \
  --max-train-samples 4096 \
  --epochs 5
  • setup_distributed_from_env() initializes the process group from torchrun environment variables and selects the local CUDA device
  • Data is sharded across ranks via deterministic block partitioning
  • Metrics are all-reduced across all processes
  • Only rank 0 writes metrics.jsonl and explicit save_*() output files
  • Cache warmup runs on all ranks with file-lock coordination

Model Contract

Your model's forward method receives two positional arguments:

Argument Shape Dtype Description
sar [B, 2, 256, 256] float32 Sentinel-1 SAR image (2 channels)
cloudy [B, 13, 256, 256] float32 Cloudy Sentinel-2 optical image (13 channels)

Output: a prediction tensor, typically [B, 13, 256, 256].

class MyModel(nn.Module):
    def forward(self, sar, cloudy):
        # sar:    [B, 2,  256, 256]
        # cloudy: [B, 13, 256, 256]
        x = torch.cat([sar, cloudy], dim=1)  # [B, 15, 256, 256]
        return self.network(x)

The loss and metric functions receive (prediction, batch) where batch is the full dict containing "sar", "cloudy", "target", and "meta":

def my_loss(prediction, batch):
    return F.l1_loss(prediction, batch["target"])

License

MIT

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages