From f0d32fa96c818425fb104cb611dec25e885bb7de Mon Sep 17 00:00:00 2001 From: Severin Magel <116261790+sevmag@users.noreply.github.com> Date: Sat, 15 Aug 2026 14:40:36 -0400 Subject: [PATCH 1/2] spine: rename the pretext package to pretrain The directory now matches what it is used for; module paths become spine.pretrain.*, including the config _target_ strings. Existing TransferCheckpoints store the old spine.pretext.* target as a string but never re-instantiate it (only backbone._target_ is read, for backbone detection), so previously trained encoders keep loading. Co-Authored-By: Claude Fable 5 --- README.md | 2 +- configs/callbacks/curtain_auc.yaml | 2 +- configs/task/curtain.yaml | 2 +- configs/task/objectives/v1.yaml | 2 +- configs/task/objectives/v2.yaml | 4 ++-- examples/graphnet_demo.py | 10 +++++----- src/spine/data/datamodule.py | 2 +- src/spine/{pretext => pretrain}/__init__.py | 0 src/spine/{pretext => pretrain}/base.py | 0 src/spine/{pretext => pretrain}/curtain/__init__.py | 0 src/spine/{pretext => pretrain}/curtain/callbacks.py | 2 +- src/spine/{pretext => pretrain}/curtain/head.py | 2 +- src/spine/{pretext => pretrain}/curtain/objectives.py | 2 +- src/spine/{pretext => pretrain}/curtain/sampler.py | 0 src/spine/{pretext => pretrain}/curtain/task.py | 6 +++--- src/spine/ssl_module.py | 2 +- src/spine/train.py | 2 +- 17 files changed, 20 insertions(+), 20 deletions(-) rename src/spine/{pretext => pretrain}/__init__.py (100%) rename src/spine/{pretext => pretrain}/base.py (100%) rename src/spine/{pretext => pretrain}/curtain/__init__.py (100%) rename src/spine/{pretext => pretrain}/curtain/callbacks.py (98%) rename src/spine/{pretext => pretrain}/curtain/head.py (99%) rename src/spine/{pretext => pretrain}/curtain/objectives.py (97%) rename src/spine/{pretext => pretrain}/curtain/sampler.py (100%) rename src/spine/{pretext => pretrain}/curtain/task.py (98%) diff --git a/README.md b/README.md index e9431d7..67887f3 100644 --- a/README.md +++ b/README.md @@ -70,7 +70,7 @@ Datasets and you keep them disjoint. Every selected event must satisfy the task's sampling requirements; tasks raise on events that fall short instead of skipping them silently, so pre filter your selection with the task's own predicate. For CURTAIN that is -`spine.pretext.curtain.sampler.can_always_split`, called with the same +`spine.pretrain.curtain.sampler.can_always_split`, called with the same `min_visible`/`min_future` you give the task and float32 times. **4. Feature scaling.** A `FeatureScaler` subclass (`scale_pulses` and diff --git a/configs/callbacks/curtain_auc.yaml b/configs/callbacks/curtain_auc.yaml index a2677cf..a4cdf94 100644 --- a/configs/callbacks/curtain_auc.yaml +++ b/configs/callbacks/curtain_auc.yaml @@ -1 +1 @@ -- _target_: spine.pretext.curtain.callbacks.CurtainValAUC +- _target_: spine.pretrain.curtain.callbacks.CurtainValAUC diff --git a/configs/task/curtain.yaml b/configs/task/curtain.yaml index a15ed7b..bc453b5 100644 --- a/configs/task/curtain.yaml +++ b/configs/task/curtain.yaml @@ -8,7 +8,7 @@ defaults: - objectives: v1 - _self_ -_target_: spine.pretext.curtain.task.CurtainTask +_target_: spine.pretrain.curtain.task.CurtainTask max_pulses: 768 center_time: true dt_scale: 500.0 diff --git a/configs/task/objectives/v1.yaml b/configs/task/objectives/v1.yaml index 9c9ce52..49eb260 100644 --- a/configs/task/objectives/v1.yaml +++ b/configs/task/objectives/v1.yaml @@ -1,4 +1,4 @@ # @package task # v1: occupancy only objectives: - - _target_: spine.pretext.curtain.objectives.OccupancyObjective + - _target_: spine.pretrain.curtain.objectives.OccupancyObjective diff --git a/configs/task/objectives/v2.yaml b/configs/task/objectives/v2.yaml index bc02feb..d366dc6 100644 --- a/configs/task/objectives/v2.yaml +++ b/configs/task/objectives/v2.yaml @@ -1,6 +1,6 @@ # @package task # v2: occupancy + Delta-t (cwm-referenced) objectives: - - _target_: spine.pretext.curtain.objectives.OccupancyObjective - - _target_: spine.pretext.curtain.objectives.DtObjective + - _target_: spine.pretrain.curtain.objectives.OccupancyObjective + - _target_: spine.pretrain.curtain.objectives.DtObjective weight: 1.0 diff --git a/examples/graphnet_demo.py b/examples/graphnet_demo.py index 3ebac10..5b6f2f1 100644 --- a/examples/graphnet_demo.py +++ b/examples/graphnet_demo.py @@ -39,10 +39,10 @@ from spine.data.geometry import load_geometry from spine.data.scaling import FeatureLayout -from spine.pretext.curtain.callbacks import CurtainValAUC -from spine.pretext.curtain.objectives import OccupancyObjective -from spine.pretext.curtain.sampler import can_always_split -from spine.pretext.curtain.task import CurtainTask +from spine.pretrain.curtain.callbacks import CurtainValAUC +from spine.pretrain.curtain.objectives import OccupancyObjective +from spine.pretrain.curtain.sampler import can_always_split +from spine.pretrain.curtain.task import CurtainTask from spine.train import fit LAYOUT = FeatureLayout() @@ -271,7 +271,7 @@ def main() -> None: task = CurtainTask( geo=geo, # v2 is one line more: append DtObjective(weight=1.0) from - # spine.pretext.curtain.objectives + # spine.pretrain.curtain.objectives objectives=[OccupancyObjective()], scaler=DetectorScaler(Prometheus(), PULSE_FEATURES), dt_scale=100.0, diff --git a/src/spine/data/datamodule.py b/src/spine/data/datamodule.py index eb77be5..2ac1849 100644 --- a/src/spine/data/datamodule.py +++ b/src/spine/data/datamodule.py @@ -19,7 +19,7 @@ import pytorch_lightning as pl from torch.utils.data import DataLoader, Dataset -from spine.pretext.base import PretextTask +from spine.pretrain.base import PretextTask class RawEvent(TypedDict): diff --git a/src/spine/pretext/__init__.py b/src/spine/pretrain/__init__.py similarity index 100% rename from src/spine/pretext/__init__.py rename to src/spine/pretrain/__init__.py diff --git a/src/spine/pretext/base.py b/src/spine/pretrain/base.py similarity index 100% rename from src/spine/pretext/base.py rename to src/spine/pretrain/base.py diff --git a/src/spine/pretext/curtain/__init__.py b/src/spine/pretrain/curtain/__init__.py similarity index 100% rename from src/spine/pretext/curtain/__init__.py rename to src/spine/pretrain/curtain/__init__.py diff --git a/src/spine/pretext/curtain/callbacks.py b/src/spine/pretrain/curtain/callbacks.py similarity index 98% rename from src/spine/pretext/curtain/callbacks.py rename to src/spine/pretrain/curtain/callbacks.py index a5a5030..aa844c4 100644 --- a/src/spine/pretext/curtain/callbacks.py +++ b/src/spine/pretrain/curtain/callbacks.py @@ -11,7 +11,7 @@ import pytorch_lightning as pl from pytorch_lightning.callbacks import Callback -from spine.pretext.curtain.task import real_query_mask +from spine.pretrain.curtain.task import real_query_mask def auc(scores: np.ndarray, labels: np.ndarray) -> float: diff --git a/src/spine/pretext/curtain/head.py b/src/spine/pretrain/curtain/head.py similarity index 99% rename from src/spine/pretext/curtain/head.py rename to src/spine/pretrain/curtain/head.py index 7284e5a..ac8f2ad 100644 --- a/src/spine/pretext/curtain/head.py +++ b/src/spine/pretrain/curtain/head.py @@ -13,7 +13,7 @@ from torch import Tensor, nn from spine.backbones.base import EncodedEvent -from spine.pretext.base import Objective +from spine.pretrain.base import Objective class PositionQueryEncoder(nn.Module): diff --git a/src/spine/pretext/curtain/objectives.py b/src/spine/pretrain/curtain/objectives.py similarity index 97% rename from src/spine/pretext/curtain/objectives.py rename to src/spine/pretrain/curtain/objectives.py index d11d630..df273c1 100644 --- a/src/spine/pretext/curtain/objectives.py +++ b/src/spine/pretrain/curtain/objectives.py @@ -9,7 +9,7 @@ import torch.nn.functional as F from torch import Tensor, nn -from spine.pretext.base import Objective +from spine.pretrain.base import Objective class OccupancyObjective(Objective): diff --git a/src/spine/pretext/curtain/sampler.py b/src/spine/pretrain/curtain/sampler.py similarity index 100% rename from src/spine/pretext/curtain/sampler.py rename to src/spine/pretrain/curtain/sampler.py diff --git a/src/spine/pretext/curtain/task.py b/src/spine/pretrain/curtain/task.py similarity index 98% rename from src/spine/pretext/curtain/task.py rename to src/spine/pretrain/curtain/task.py index bdf0938..4351bd0 100644 --- a/src/spine/pretext/curtain/task.py +++ b/src/spine/pretrain/curtain/task.py @@ -12,9 +12,9 @@ from torch import Tensor, nn from spine.data.scaling import FeatureScaler -from spine.pretext.base import Objective, PretextTask, Sample -from spine.pretext.curtain.head import MultiObjectiveHead -from spine.pretext.curtain.sampler import sample_event +from spine.pretrain.base import Objective, PretextTask, Sample +from spine.pretrain.curtain.head import MultiObjectiveHead +from spine.pretrain.curtain.sampler import sample_event def real_query_mask(pred: Tensor, batch: dict) -> Tensor: diff --git a/src/spine/ssl_module.py b/src/spine/ssl_module.py index 9e8acdb..5deaf17 100644 --- a/src/spine/ssl_module.py +++ b/src/spine/ssl_module.py @@ -13,7 +13,7 @@ from torch import nn from spine.backbones.base import Backbone -from spine.pretext.base import PretextTask +from spine.pretrain.base import PretextTask class SSLModule(pl.LightningModule): diff --git a/src/spine/train.py b/src/spine/train.py index ef3e4e1..498a440 100644 --- a/src/spine/train.py +++ b/src/spine/train.py @@ -18,7 +18,7 @@ from spine.backbones.base import Backbone from spine.data.datamodule import SpineDataModule -from spine.pretext.base import PretextTask +from spine.pretrain.base import PretextTask from spine.ssl_module import SSLModule from spine.utils import TransferCheckpoint From ca5f90fef3fd0c1ca84b248dae520d2c12f02671 Mon Sep 17 00:00:00 2001 From: Severin Magel <116261790+sevmag@users.noreply.github.com> Date: Sat, 15 Aug 2026 14:41:52 -0400 Subject: [PATCH 2/2] spine: rename the Pretext identifiers to Pretrain Follows the package move: PretextTask -> PretrainTask, PretextDataset -> PretrainDataset, and the surrounding prose. This is a public API change -- downstream code importing PretextTask must be updated with it. Co-Authored-By: Claude Fable 5 --- DESIGN.md | 20 +++++++++--------- README.md | 12 +++++------ configs/task/curtain.yaml | 2 +- configs/trainer/default.yaml | 2 +- src/spine/backbones/base.py | 2 +- src/spine/data/__init__.py | 2 +- src/spine/data/datamodule.py | 28 ++++++++++++------------- src/spine/data/scaling.py | 2 +- src/spine/pretrain/__init__.py | 2 +- src/spine/pretrain/base.py | 6 +++--- src/spine/pretrain/curtain/__init__.py | 2 +- src/spine/pretrain/curtain/callbacks.py | 2 +- src/spine/pretrain/curtain/sampler.py | 2 +- src/spine/pretrain/curtain/task.py | 4 ++-- src/spine/ssl_module.py | 12 +++++------ src/spine/train.py | 6 +++--- src/spine/utils.py | 2 +- 17 files changed, 54 insertions(+), 54 deletions(-) diff --git a/DESIGN.md b/DESIGN.md index da2351d..0c63c3c 100644 --- a/DESIGN.md +++ b/DESIGN.md @@ -2,27 +2,27 @@ Self-supervised Pretraining In Neutrino Experiments. The repo produces **pretrained backbones** (encoder checkpoints) that downstream supervised -benchmarks fine-tune. First pretext: **CURTAIN** (occupancy / light-front +benchmarks fine-tune. First pretrain: **CURTAIN** (occupancy / light-front forecast); built so new SSL methods are small plugins. ## Mental model -A run is: **Data → Backbone → Pretext(head + targets + loss)**, wired by an +A run is: **Data → Backbone → Pretrain(head + targets + loss)**, wired by an **Engine**, named by a **Config**. The only thing you write to add a method is a -`pretext/` plugin — data, backbone, and engine are reused unchanged. +`pretrain/` plugin — data, backbone, and engine are reused unchanged. ## Blocks | block | responsibility | |---|---| | `data/` | geometry asset + sensor-key lookup, FeatureScaler scaling, datamodule (reader- & selection-agnostic) | | `backbones/` | encoder interface (swappable; graphnet-free; DeepIce impl in integrations/spine_graphnet/) | -| `pretext/` | pretext-task interface + `curtain/` (sampler, head, objectives, task, val callbacks) | +| `pretrain/` | pretrain-task interface + `curtain/` (sampler, head, objectives, task, val callbacks) | | `ssl_module.py` | Lightning module; optimizer/scheduler injected as factories (transfer export in `utils.py`) | | `configs/` + `train.py` | Hydra groups compose a run (examples/train_curtain.py); fit() assembles | ## The two interfaces (all extensibility lives here) - **`Backbone.encode(batch) -> EncodedEvent(tokens, token_mask, cls)`** — swap - architectures without touching pretext/engine. -- **`PretextTask`** — `make_sample` (CPU: mask/target), `collate`, `build_head`, + architectures without touching pretrain/engine. +- **`PretrainTask`** — `make_sample` (CPU: mask/target), `collate`, `build_head`, `loss`. A task carries a list of weighted **`Objective`s**, each an abstract class owning its own head (`build_head`) and `loss`, over one sample. @@ -44,7 +44,7 @@ model/dataset/train files. `ckpt["backbone"]` into graphnet DeepIce, so the exported state_dict must stay compatible — keep DeepIce, or vendor a state-dict-identical encoder later (`examples/deepice_backbone.py` TODO). -- **Data layer.** Pretext needs **raw** pulses (the Δt reference is +- **Data layer.** Pretrain needs **raw** pulses (the Δt reference is charge-weighted-mean-time on raw values), so standardization runs at the model boundary **after** the split, not in the source. LMDB is welcome for speed but as the **low-level read utilities** behind the read `Dataset` (raw pulses; identity @@ -100,8 +100,8 @@ best val (rank-0 only). Downstream loads `ckpt["backbone"]`. Finetuning/eval stays in the existing bench — this repo emits encoders, nothing more. ## Adding a method (extensibility test) -New folder under `pretext/`, point a `task/.yaml` `_target_` at the -new `PretextTask`: +New folder under `pretrain/`, point a `task/.yaml` `_target_` at the +new `PretrainTask`: - **MAE**: `make_sample` masks pulses; head = decoder; loss = reconstruct. - **Contrastive**: `make_sample` = two views; head = projection on `cls`; loss = NT-Xent. Data, backbone, engine unchanged. @@ -111,7 +111,7 @@ Data, backbone, engine unchanged. reference pretraining (best val loss within noise, AUCs within 7e-4). 2. Reproduce v2 by config (`task/objectives=v2` exists; revalidation open). 3. Config system (hydra) ✓ — profile loader remains. -Deferred: other backbones, other pretexts, multi-detector, in-repo eval. +Deferred: other backbones, other pretraining tasks, multi-detector, in-repo eval. ## Open decisions 1. graphnet DeepIce behind the interface vs **vendor** a standalone encoder. diff --git a/README.md b/README.md index 67887f3..c3814be 100644 --- a/README.md +++ b/README.md @@ -5,13 +5,13 @@ data. The repo produces **pretrained backbones** (encoder checkpoints) that downstream supervised benchmarks load and fine-tune. A *spine* is a backbone, which is exactly what this emits. -The first pretext task is **CURTAIN** (occupancy / light-front forecast). The +The first pretrain task is **CURTAIN** (occupancy / light-front forecast). The architecture is built so a new self-supervised method is a small plugin under -`pretext/`, reusing the data, backbone, and training engine unchanged. +`pretrain/`, reusing the data, backbone, and training engine unchanged. ```mermaid flowchart LR - D["your data
reader + geometry"] --> T["PretextTask
e.g. CURTAIN"] + D["your data
reader + geometry"] --> T["PretrainTask
e.g. CURTAIN"] T --> B["Backbone
e.g. DeepIce"] B --> H["objective heads
+ loss"] H --> X["pretrained backbone
for your fine tune"] @@ -21,7 +21,7 @@ flowchart LR The core is framework-agnostic and fits neatly into plain PyTorch: it depends only on torch, pytorch-lightning and numpy. Readers are ordinary indexable `Dataset`s emitting a small canonical sample format, models are `nn.Module`s -behind two narrow interfaces (`Backbone`, `PretextTask`), and `fit()` takes +behind two narrow interfaces (`Backbone`, `PretrainTask`), and `fit()` takes injected factories and callbacks. Hydra and graphnet integrate neatly, but both are strictly optional conveniences: use either, both, or neither. Around that core you choose your frame: @@ -43,7 +43,7 @@ infrastructure, splits, logging, versioning) stays yours. **1. Raw events.** Any PyTorch `Dataset` yielding `raw[i] -> {"event_no": int, "pulses": [P, F] float32, "sensor_key": [P] int}` -(stated canonically in `spine/data/datamodule.py`). Pulses stay raw: pretext +(stated canonically in `spine/data/datamodule.py`). Pulses stay raw: pretrain tasks make their sampling decisions and build their targets in detector units, and standardization happens later at collate. Columns follow the task's `FeatureLayout`, by default `(x, y, z, t, charge)`; pass a different layout @@ -90,7 +90,7 @@ a jagged NJT. src/spine/ data/ geometry + FeatureScaler scaling, datamodule (reader- & selection-agnostic) backbones/ encoder interface (swappable; graphnet-free core) - pretext/ pretext-task interface + curtain/ (the first task) + pretrain/ pretrain-task interface + curtain/ (the first task) ssl_module.py Lightning SSLModule (optimizer/scheduler injected as factories) utils.py TransferCheckpoint callback (best-val backbone export) train.py reader-agnostic fit() assembly diff --git a/configs/task/curtain.yaml b/configs/task/curtain.yaml index bc453b5..532fae4 100644 --- a/configs/task/curtain.yaml +++ b/configs/task/curtain.yaml @@ -1,4 +1,4 @@ -# The CURTAIN pretext task -- a self-contained plugin config that pulls in its +# The CURTAIN pretrain task -- a self-contained plugin config that pulls in its # own objectives. The launcher injects the runtime geo + scaler and # recursive-instantiates the rest. Sampler knobs (q_lo/q_hi, # pos_k, neg_anchor, rand_neg_frac, min_visible, min_future, resample_tries) diff --git a/configs/trainer/default.yaml b/configs/trainer/default.yaml index 4e206eb..2a0cf65 100644 --- a/configs/trainer/default.yaml +++ b/configs/trainer/default.yaml @@ -2,7 +2,7 @@ batch: 64 num_workers: 16 val_num_workers: null # null -> num_workers devices: 1 -precision: "32-true" # bf16-mixed is faster but can degrade val post-plateau for this pretext +precision: "32-true" # bf16-mixed is faster but can degrade val post-plateau for this pretrain max_epochs: 200 patience: 15 # EarlyStopping (epochs) grad_clip: 1.0 diff --git a/src/spine/backbones/base.py b/src/spine/backbones/base.py index 3c21cf3..5a78b03 100644 --- a/src/spine/backbones/base.py +++ b/src/spine/backbones/base.py @@ -1,6 +1,6 @@ """Backbone interface: a collated batch -> per-token embeddings + CLS. -Swapping encoders means implementing `encode`; pretext and engine code stay +Swapping encoders means implementing `encode`; pretrain and engine code stay unchanged. """ diff --git a/src/spine/data/__init__.py b/src/spine/data/__init__.py index 15c8704..b4f6648 100644 --- a/src/spine/data/__init__.py +++ b/src/spine/data/__init__.py @@ -1 +1 @@ -"""Data layer: pretext datamodule, geometry asset, feature scaling.""" +"""Data layer: pretrain datamodule, geometry asset, feature scaling.""" diff --git a/src/spine/data/datamodule.py b/src/spine/data/datamodule.py index 2ac1849..efc8616 100644 --- a/src/spine/data/datamodule.py +++ b/src/spine/data/datamodule.py @@ -1,12 +1,12 @@ -"""Raw-pulse read Datasets -> pretext samples. +"""Raw-pulse read Datasets -> pretrain samples. THE read contract: raw[i] -> {"event_no": int, "pulses": [P, F] raw, "sensor_key": [P] int} -- feature columns per the task's FeatureLayout, raw -values (standardization happens after the pretext split), sensor keys matching +values (standardization happens after the pretrain split), sensor keys matching the geometry asset's key array (multi-level IDs composed by the reader; single-PMT detectors use 1 for the missing level). SPINE ships a minimal reference reader (spine.data.readers); graphnet-backed readers live in -spine_graphnet. PretextDataset is a pure index -> sample map -- +spine_graphnet. PretrainDataset is a pure index -> sample map -- make_sample raises on events it cannot use, so batches are never silently short (an empty batch deadlocks DDP). """ @@ -19,7 +19,7 @@ import pytorch_lightning as pl from torch.utils.data import DataLoader, Dataset -from spine.pretrain.base import PretextTask +from spine.pretrain.base import PretrainTask class RawEvent(TypedDict): @@ -38,20 +38,20 @@ def __len__(self) -> int: ... def __getitem__(self, idx: int) -> RawEvent: ... -class PretextDataset(Dataset): - """Transform on top of a read Dataset: index -> pretext sample.""" +class PretrainDataset(Dataset): + """Transform on top of a read Dataset: index -> pretrain sample.""" def __init__( self, raw: RawPulseDataset, - task: PretextTask, + task: PretrainTask, resample: bool = True, ): - """Compose the pretext transform over a read Dataset. + """Compose the pretrain transform over a read Dataset. Args: raw: Read-layer Dataset satisfying the RawPulseDataset contract. - task: Pretext task whose make_sample transforms each event. + task: Pretrain task whose make_sample transforms each event. resample: Fresh RNG per call (training) instead of a fixed per-index seed (validation). """ @@ -76,13 +76,13 @@ def __getitem__(self, idx: int): class SpineDataModule(pl.LightningDataModule): - """Train/val DataLoaders over PretextDataset with the task's collate.""" + """Train/val DataLoaders over PretrainDataset with the task's collate.""" def __init__( self, train_raw: RawPulseDataset, val_raw: RawPulseDataset, - task: PretextTask, + task: PretrainTask, batch_size: int = 64, num_workers: int = 16, val_num_workers: int | None = None, @@ -92,7 +92,7 @@ def __init__( Args: train_raw: Read Dataset for the training events. val_raw: Read Dataset for the validation events. - task: Pretext task providing make_sample and collate. + task: Pretrain task providing make_sample and collate. batch_size: Events per batch for both loaders. num_workers: Worker processes for the training loader. val_num_workers: Worker processes for the validation loader; @@ -115,7 +115,7 @@ def _loader( # runs NCCL/CUDA threads can deadlock a DDP rank; spawn children start # clean, and persistent workers pay the startup cost once. return DataLoader( - PretextDataset(raw, self.task, resample=resample), + PretrainDataset(raw, self.task, resample=resample), batch_size=self.batch_size, shuffle=shuffle, num_workers=workers, @@ -126,7 +126,7 @@ def _loader( ) def train_dataloader(self) -> DataLoader: - """Shuffled drop-last loader; a fresh pretext split every epoch. + """Shuffled drop-last loader; a fresh pretrain split every epoch. Returns: The training DataLoader. diff --git a/src/spine/data/scaling.py b/src/spine/data/scaling.py index e2dbad2..a3980b8 100644 --- a/src/spine/data/scaling.py +++ b/src/spine/data/scaling.py @@ -65,7 +65,7 @@ def scale_pulses(self, x: Tensor) -> Tensor: @abstractmethod def scale_positions(self, p: Tensor) -> Tensor: - """Standardize raw positions (pretext query coordinates). + """Standardize raw positions (pretrain query coordinates). Args: p: [..., 3] raw positions, same units as the pulse xyz columns. diff --git a/src/spine/pretrain/__init__.py b/src/spine/pretrain/__init__.py index cb63398..9e15876 100644 --- a/src/spine/pretrain/__init__.py +++ b/src/spine/pretrain/__init__.py @@ -1 +1 @@ -"""Pretext-task interface and task implementations.""" +"""Pretrain-task interface and task implementations.""" diff --git a/src/spine/pretrain/base.py b/src/spine/pretrain/base.py index 814cce3..e94657e 100644 --- a/src/spine/pretrain/base.py +++ b/src/spine/pretrain/base.py @@ -1,4 +1,4 @@ -"""Pretext-task interface: make_sample -> collate -> build_head -> loss. +"""Pretrain-task interface: make_sample -> collate -> build_head -> loss. A task owns its per-event sampling, batching, head construction and loss; `Objective`s are its weighted sub-targets, each bringing its own head and @@ -64,7 +64,7 @@ def loss(self, pred: Tensor, batch: dict) -> Tensor: ... -class PretextTask(ABC): +class PretrainTask(ABC): """Factory + transform + loss for one self-supervised objective.""" #: objectives this task scores (defines head width and the loss terms) @@ -78,7 +78,7 @@ def make_sample( Args: event: One raw event from the read layer. - rng: Per-call generator; fresh entropy resamples the pretext, + rng: Per-call generator; fresh entropy resamples the pretrain, a fixed seed reproduces it. Returns: diff --git a/src/spine/pretrain/curtain/__init__.py b/src/spine/pretrain/curtain/__init__.py index 04de8f0..076b6a0 100644 --- a/src/spine/pretrain/curtain/__init__.py +++ b/src/spine/pretrain/curtain/__init__.py @@ -1 +1 @@ -"""CURTAIN: the occupancy / light-front forecast pretext.""" +"""CURTAIN: the occupancy / light-front forecast pretrain.""" diff --git a/src/spine/pretrain/curtain/callbacks.py b/src/spine/pretrain/curtain/callbacks.py index aa844c4..15351ba 100644 --- a/src/spine/pretrain/curtain/callbacks.py +++ b/src/spine/pretrain/curtain/callbacks.py @@ -1,4 +1,4 @@ -"""Validation callbacks for the CURTAIN pretext. +"""Validation callbacks for the CURTAIN pretrain. Epoch-global metrics (AUC is rank-based over the full val set) cannot flow through per-batch log averaging, so callbacks cache per batch and reduce once diff --git a/src/spine/pretrain/curtain/sampler.py b/src/spine/pretrain/curtain/sampler.py index 0382955..40c0ca1 100644 --- a/src/spine/pretrain/curtain/sampler.py +++ b/src/spine/pretrain/curtain/sampler.py @@ -131,7 +131,7 @@ def sample_event( min_future: int, resample_tries: int, ) -> dict | None: - """Build the pretext split for one event. + """Build the pretrain split for one event. Args: pt: Pulse times. diff --git a/src/spine/pretrain/curtain/task.py b/src/spine/pretrain/curtain/task.py index 4351bd0..8dd57bb 100644 --- a/src/spine/pretrain/curtain/task.py +++ b/src/spine/pretrain/curtain/task.py @@ -12,7 +12,7 @@ from torch import Tensor, nn from spine.data.scaling import FeatureScaler -from spine.pretrain.base import Objective, PretextTask, Sample +from spine.pretrain.base import Objective, PretrainTask, Sample from spine.pretrain.curtain.head import MultiObjectiveHead from spine.pretrain.curtain.sampler import sample_event @@ -34,7 +34,7 @@ def real_query_mask(pred: Tensor, batch: dict) -> Tensor: return torch.arange(pred.shape[1], device=pred.device)[None] < qlen[:, None] -class CurtainTask(PretextTask): +class CurtainTask(PretrainTask): """Occupancy(/+dt) forecast over held-out sensors of a split event.""" def __init__( diff --git a/src/spine/ssl_module.py b/src/spine/ssl_module.py index 5deaf17..7c49066 100644 --- a/src/spine/ssl_module.py +++ b/src/spine/ssl_module.py @@ -1,4 +1,4 @@ -"""Lightning module wiring a Backbone to a PretextTask's head. +"""Lightning module wiring a Backbone to a PretrainTask's head. Logs with sync_dist=True so ReduceLROnPlateau steps identically on every DDP rank; exposes `backbone` and `model` for TransferCheckpoint. Batch access @@ -13,16 +13,16 @@ from torch import nn from spine.backbones.base import Backbone -from spine.pretrain.base import PretextTask +from spine.pretrain.base import PretrainTask class SSLModule(pl.LightningModule): - """Backbone + pretext head; the task computes the loss.""" + """Backbone + pretrain head; the task computes the loss.""" def __init__( self, backbone: Backbone, - task: PretextTask, + task: PretrainTask, optimizer: Callable, scheduler: Callable | None = None, scheduler_config: dict | None = None, @@ -31,7 +31,7 @@ def __init__( Args: backbone: Encoder to pretrain. - task: Pretext task; builds the head and computes the loss. + task: Pretrain task; builds the head and computes the loss. optimizer: Factory, parameters -> torch Optimizer (partial). scheduler: Optional factory, optimizer -> LR scheduler. scheduler_config: Lightning lr_scheduler metadata; None means @@ -58,7 +58,7 @@ def backbone(self) -> nn.Module: @property def head(self) -> nn.Module: - """The pretext prediction head.""" + """The pretrain prediction head.""" return self.model["head"] def forward(self, batch: dict): diff --git a/src/spine/train.py b/src/spine/train.py index 498a440..d30aab0 100644 --- a/src/spine/train.py +++ b/src/spine/train.py @@ -18,7 +18,7 @@ from spine.backbones.base import Backbone from spine.data.datamodule import SpineDataModule -from spine.pretrain.base import PretextTask +from spine.pretrain.base import PretrainTask from spine.ssl_module import SSLModule from spine.utils import TransferCheckpoint @@ -26,7 +26,7 @@ def fit( train_raw: Dataset, val_raw: Dataset, - task: PretextTask, + task: PretrainTask, backbone: Backbone, out: str, *, @@ -50,7 +50,7 @@ def fit( Args: train_raw: Read Dataset for the training events. val_raw: Read Dataset for the validation events. - task: Pretext task (sampling, collate, head, loss). + task: Pretrain task (sampling, collate, head, loss). backbone: Encoder to pretrain; its state_dict is the exported artifact. out: Path the transfer checkpoint is written to on best val loss. optimizer: Factory mapping parameters -> a torch Optimizer. diff --git a/src/spine/utils.py b/src/spine/utils.py index af55bee..21d83c7 100644 --- a/src/spine/utils.py +++ b/src/spine/utils.py @@ -13,7 +13,7 @@ class TransferCheckpoint(Callback): """Save the backbone (+ full module) when `val_loss_epoch` improves. Only rank 0 writes under DDP; reads `pl_module.backbone` (the exported - encoder) and `pl_module.model` (the full pretext model). + encoder) and `pl_module.model` (the full pretrain model). """ def __init__(self, out: str, config: dict, min_delta: float = 1e-4):