One-class training toolkit for satellite cloud removal -- deterministic sampling, fast HF block streaming, and DDP out of the box.
English | 한국어
Table of Contents
- Single-class API --
Trainerwithstep()+test(), nothing else to learn - Deterministic block sampling -- one-seed system (
seed) for exact reproducibility - HF v2 block data -- streams compressed
.crpackblocks 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
uv add git+https://github.com/smturtle2/cr-train.gitfrom 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
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-examplePass --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.
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 9Output shows the raw draw order, the final selected block indices, and a bitmap of selected (■) vs. skipped (□) logical blocks.
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.jsonlin 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.
| 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. |
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,
}Runs test evaluation with the current model state. Returns:
{
"epoch": 2,
"loss": 0.0387,
"metrics": {"mae": 0.0295},
"num_samples": 256,
"num_batches": 64,
}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.
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,
}Writes model weights only. When path is omitted, the file is written to
<output_dir>/model-epoch-XXXX.pt.
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().
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.
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,
}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_sampleSee examples/bitmask_sampling_demo.py for a full visualization.
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
seedvalues 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 oncenum_workers > 0. - Persistent local datasets are never auto-deleted. Remove
dataset_dirmanually to reclaim disk space.
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 5setup_distributed_from_env()initializes the process group fromtorchrunenvironment 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.jsonland explicitsave_*()output files - Cache warmup runs on all ranks with file-lock coordination
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"])