Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
105 commits
Select commit Hold shift + click to select a range
4d27859
Add reg module
georgeyiasemis Jul 22, 2024
4700070
Update reg and add integration
georgeyiasemis Jul 25, 2024
0643b1e
Add smooth loss
georgeyiasemis Jul 25, 2024
dc1655e
Update transforms
georgeyiasemis Jul 25, 2024
89d3e62
Minor fix
georgeyiasemis Jul 25, 2024
952b710
Minor fix
georgeyiasemis Jul 25, 2024
f555921
Add option to save mask
georgeyiasemis Aug 1, 2024
e0a1eef
Minor fix
georgeyiasemis Aug 1, 2024
a4be1d9
Add recon and eval
georgeyiasemis Aug 1, 2024
741e700
Add adpt
georgeyiasemis Aug 1, 2024
8020099
Option to use acs as mask
georgeyiasemis Aug 5, 2024
e319efc
Forgotten return
georgeyiasemis Aug 5, 2024
f468869
Recover if error
georgeyiasemis Aug 5, 2024
8742423
Visualize registration items
georgeyiasemis Aug 12, 2024
cbf5a56
Return registered image if registration
georgeyiasemis Aug 13, 2024
2eaed0b
Add traditional recon as models
georgeyiasemis Aug 21, 2024
6392c3c
Choose if calculate displacement
georgeyiasemis Aug 21, 2024
e840963
Minor fix
georgeyiasemis Aug 21, 2024
f75f8e8
Minor fix
georgeyiasemis Aug 21, 2024
d912b9e
Minor fix
georgeyiasemis Aug 21, 2024
d7a2975
Add 3d models
Sep 21, 2024
64e0345
Add 3d varnet and 3d kspace varnet
Sep 21, 2024
8a7d731
Minor fixes
georgeyiasemis Sep 21, 2024
6f52a4b
Update for registration
georgeyiasemis Sep 21, 2024
10124c5
Add vit transformers module
georgeyiasemis Oct 20, 2024
842dd87
Rename
georgeyiasemis Oct 20, 2024
8d7cd07
Minor fix
georgeyiasemis Oct 20, 2024
d7ca4a5
Minor fix
georgeyiasemis Oct 20, 2024
d14e00f
Isort
georgeyiasemis Oct 20, 2024
930e2e5
Minor fix
georgeyiasemis Oct 20, 2024
a96907a
Minor fix
georgeyiasemis Oct 20, 2024
9a45371
Rename
georgeyiasemis Oct 20, 2024
5c2fea3
No torchvision version
georgeyiasemis Oct 24, 2024
adb18fd
Allow for non end to end training
georgeyiasemis Oct 24, 2024
b16eff6
Add voxelmorph model
georgeyiasemis Nov 4, 2024
5ca113b
isort, add vit reg
georgeyiasemis Nov 4, 2024
f0a2b4f
Minor fix
georgeyiasemis Nov 4, 2024
86fd9f4
Changes for saving registered images, minor fixes
georgeyiasemis Nov 21, 2024
0a8af40
Allow decoupled training if registration
georgeyiasemis Jan 28, 2025
d49a24b
Minor fix
georgeyiasemis Jan 28, 2025
29d216a
Add MEDL model
georgeyiasemis Feb 7, 2025
b96b184
Minor fic
georgeyiasemis Feb 7, 2025
b1d3fba
Minor fix
georgeyiasemis Feb 7, 2025
9773d44
feat: integrate adaptive sampling and end-to-end registration from main
georgeyiasemis Aug 5, 2026
8d729f2
fix: normalize paper experiment configs for main schema
georgeyiasemis Aug 5, 2026
e2742aa
fix: drop obsolete keys from adpt experiment dumps
georgeyiasemis Aug 5, 2026
9fd2e3d
fix: restore optimizer eps and finish adpt config cleanup
georgeyiasemis Aug 5, 2026
217468c
fix: sanitize remaining adpt masking keys for MaskFuncMode API
georgeyiasemis Aug 5, 2026
6ac9c3a
fix: rebuild broken inference transforms in adpt config dumps
georgeyiasemis Aug 5, 2026
3fda588
fix: rewrite inference blocks in remaining adpt config dumps
georgeyiasemis Aug 5, 2026
27be1b2
fix: drop VarNet k-space and harden adaptive MPS/smoke paths
georgeyiasemis Aug 5, 2026
3643c38
chore: remove mini experiment configs from the branch
georgeyiasemis Aug 5, 2026
10581c8
chore: flatten paper experiment configs to named yaml files
georgeyiasemis Aug 6, 2026
226ac4c
chore: remove mini configs reintroduced by flatten commit
georgeyiasemis Aug 6, 2026
fb97825
chore: drop no_df_loss from e2e_ads_reg experiment config names
georgeyiasemis Aug 6, 2026
f0a179c
chore: remove mini configs accidentally re-added
georgeyiasemis Aug 6, 2026
a7c0572
chore: gitignore local e2e_ads mini configs
georgeyiasemis Aug 6, 2026
8e824b9
chore: remove and gitignore local e2e_ads smoke configs
georgeyiasemis Aug 6, 2026
f4a9cd7
chore: remove local mini/smoke paths from .gitignore
georgeyiasemis Aug 6, 2026
94643d3
fix: drop training.eps from config and experiment dumps
georgeyiasemis Aug 6, 2026
f7e3923
chore: require Python 3.12 and build elasticdeform from source
georgeyiasemis Aug 6, 2026
d05fb2e
docs: modernize adaptive sampling package docs and typing
georgeyiasemis Aug 6, 2026
7471105
docs: modernize registration packages and lazy-load elasticdeform
georgeyiasemis Aug 6, 2026
0f7e3bd
docs: modernize MEDL docs and switch to torch.amp autocast
georgeyiasemis Aug 6, 2026
37fb806
feat: add key-based losses and displacement-field TensorBoard viz
georgeyiasemis Aug 6, 2026
4cd87cc
fix: catch adaptive rejection sampling via typed exception
georgeyiasemis Aug 6, 2026
2a7b6b1
chore: drop empty inference.metrics dump noise from experiment configs
georgeyiasemis Aug 6, 2026
7f67336
style: reformat e2e_ads experiment YAMLs with 4-space indent
georgeyiasemis Aug 6, 2026
4d02d24
chore: remove df_0_01 displacement-loss ablation configs
georgeyiasemis Aug 6, 2026
219f31b
chore: remove df_loss displacement-loss ablation configs
georgeyiasemis Aug 6, 2026
643975d
chore: rename e2e_ads experiment configs to a concise scheme
georgeyiasemis Aug 6, 2026
b98e50a
chore: drop accidentally committed e2e_ads mini/smoke configs
georgeyiasemis Aug 6, 2026
5e722c6
fix: handle float adaptive masks in DC fill and masking
georgeyiasemis Aug 8, 2026
b02c892
fix: dynamic init masks and normalize reference_kspace
georgeyiasemis Aug 8, 2026
ed394cc
fix: adaptive policy budgeting for static ACS with dynamic k-space
georgeyiasemis Aug 8, 2026
c24e69e
fix: load paper checkpoints after ModConv bias defaults
georgeyiasemis Aug 8, 2026
61d96ab
test: unpack extra evaluate visualization return from MRI engine
georgeyiasemis Aug 8, 2026
316dda5
chore: keep validated paper YAMLs; rename e2e_ads_recon_reg
georgeyiasemis Aug 8, 2026
834a92c
docs: expand e2e_ads_recon(_reg) READMEs with paper method figures
georgeyiasemis Aug 8, 2026
46ea2e6
docs: retitle e2e_ads READMEs to paper names and clean usage notes
georgeyiasemis Aug 8, 2026
340d602
docs: fix broken RST titles and tables in e2e_ads READMEs
georgeyiasemis Aug 8, 2026
22c0a50
chore: rename dyn configs to frame-/phase-specific per papers
georgeyiasemis Aug 8, 2026
a7c16ac
chore: drop accidental kosmos_ref_predict artifacts from rename
georgeyiasemis Aug 8, 2026
f50f9aa
chore: stop tracking e2e_ads project .gitignore files
georgeyiasemis Aug 8, 2026
8795af4
revert: restore torch.where in apply_mask
georgeyiasemis Aug 8, 2026
aac5c74
refactor: remove unused centered weight norm convolutions
georgeyiasemis Aug 8, 2026
58fafc9
fix: pass time dim to DYNAMIC CreateSamplingMask
georgeyiasemis Aug 8, 2026
c545506
feat: make training log flush interval configurable
georgeyiasemis Aug 8, 2026
799391e
Merge branch 'main' into feature/adaptive-registration
georgeyiasemis Aug 8, 2026
f17cb2f
Add sibling inference YAMLs for e2e ADS projects with commented rates.
georgeyiasemis Aug 9, 2026
25a71c1
Support inference-only configs and harden model/config wiring.
georgeyiasemis Aug 9, 2026
c163b7f
Remove .hf_staging from shared .gitignore; keep it local only.
georgeyiasemis Aug 9, 2026
e918866
Drop local scratch path ignores from shared .gitignore.
georgeyiasemis Aug 9, 2026
5d633c8
Merge branch main into feature/adaptive-registration.
georgeyiasemis Aug 9, 2026
0e63f63
fix: silence ty errors and add tests for adaptive/registration modules.
georgeyiasemis Aug 9, 2026
d9f3896
test: expand coverage for adaptive, MEDL, and registration modules.
georgeyiasemis Aug 9, 2026
93f826a
fix: address Codacy and ruff findings in adaptive/registration paths.
georgeyiasemis Aug 9, 2026
551185d
test: raise coverage on writers, engines, VarNet3D, and adaptive utils.
georgeyiasemis Aug 10, 2026
dec79dc
chore: keep minimize_inference_yaml.py local-only.
georgeyiasemis Aug 10, 2026
2a7389f
chore: drop minimize_inference_yaml.py from .gitignore.
georgeyiasemis Aug 10, 2026
5efb96a
chore: drop redundant engine_name: null from project YAMLs.
georgeyiasemis Aug 10, 2026
eadb1fb
feat: visualize and save ADS initial vs predicted masks.
georgeyiasemis Aug 10, 2026
f6c0e3e
feat: print ASCII DIRECT logo when logging starts.
georgeyiasemis Aug 10, 2026
5f9b593
fix: print ASCII logo after the clinical-use warning.
georgeyiasemis Aug 10, 2026
7d25ab2
chore: update ASCII DIRECT logo art.
georgeyiasemis Aug 10, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion direct/checkpointer.py
Original file line number Diff line number Diff line change
Expand Up @@ -177,7 +177,9 @@ def _load_model(self, obj, state_dict):
# Link has more elaborate checking for incompatibles in _log_incompatible_keys
incompatible = obj.load_state_dict(state_dict, strict=False)
if incompatible.missing_keys:
raise NotImplementedError
raise RuntimeError(
f"Missing keys when loading checkpoint into model {type(obj).__name__}: {incompatible.missing_keys}"
)
if incompatible.unexpected_keys:
self.logger.warning("Unexpected keys provided which cannot be loaded: %s.", incompatible.unexpected_keys)

Expand Down
8 changes: 5 additions & 3 deletions direct/common/subsample.py
Original file line number Diff line number Diff line change
Expand Up @@ -1418,10 +1418,12 @@ def __load_masks(self, acceleration: int) -> dict[tuple[int, int], tuple[np.ndar
If the download fails.
"""
masks_path = DIRECT_CACHE_DIR / "calgary_campinas_masks"
# Accelerations are stored as floats in BaseMaskFunc; filenames use integer R.
accel = int(acceleration)
paths = [
f"R{acceleration}_218x170.npy",
f"R{acceleration}_218x174.npy",
f"R{acceleration}_218x180.npy",
f"R{accel}_218x170.npy",
f"R{accel}_218x174.npy",
f"R{accel}_218x180.npy",
]

downloaded = [download_url(self.BASE_URL + _, masks_path, md5=self.MASK_MD5S[_]) is None for _ in paths]
Expand Down
14 changes: 12 additions & 2 deletions direct/config/defaults.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,13 +28,20 @@ class TensorboardConfig(BaseConfig):
@dataclass
class LoggingConfig(BaseConfig):
log_as_image: list[str] | None = None
# How often (in iterations) to flush scalars / write TensorBoard. Default: 20.
log_interval: int = 20
tensorboard: TensorboardConfig = field(default_factory=TensorboardConfig)


@dataclass
class FunctionConfig(BaseConfig):
function: str = MISSING
multiplier: float = 1.0
# Optional tensor keys for loss comparison. When omitted, defaults are inferred
# from ``function`` (image → output_image/target, kspace → output_kspace/kspace,
# displacement_field → displacement_field/displacement_field).
source_key: str | None = None
target_key: str | None = None


@dataclass
Expand Down Expand Up @@ -106,6 +113,7 @@ class ValidationConfig(BaseConfig):
class InferenceConfig(BaseConfig):
dataset: DatasetConfig = field(default_factory=DatasetConfig)
batch_size: int = 1
metrics: list[str] = field(default_factory=list)
crop: str | None = None


Expand All @@ -130,8 +138,10 @@ class DefaultConfig(BaseConfig):

physics: PhysicsConfig = field(default_factory=PhysicsConfig)

training: TrainingConfig = field(default_factory=TrainingConfig) # This should be optional.
validation: ValidationConfig = field(default_factory=ValidationConfig) # This should be optional.
# Optional so inference-only YAMLs need not declare training/validation.
training: TrainingConfig | None = None
validation: ValidationConfig | None = None

inference: InferenceConfig | None = None

logging: LoggingConfig = field(default_factory=LoggingConfig)
64 changes: 43 additions & 21 deletions direct/data/datasets.py
Original file line number Diff line number Diff line change
Expand Up @@ -122,8 +122,11 @@ class FakeMRIBlobsDataset(Dataset):
Pass the attributes of the generated sample.
text_description: str
Description of dataset, can be useful for logging.
kspace_context: bool
If true corresponds to 3D reconstruction, else reconstruction is 2D.
kspace_context: bool or str
If set (e.g. ``True`` or ``"time"``), each item is a full 3D / 2D+time
volume of shape ``(time_or_slice, coils, height, width)``. Otherwise
reconstruction is 2D (with optional per-slice indexing when
``spatial_shape`` is 3D).
"""

def __init__(
Expand All @@ -136,7 +139,7 @@ def __init__(
filenames: list[str] | str | None = None,
pass_attrs: bool | None = None,
text_description: str | None = None,
kspace_context: bool | None = None,
kspace_context: bool | str | int | None = None,
**kwargs,
) -> None:
"""Inits :class:`FakeMRIBlobsDataset`."""
Expand All @@ -157,8 +160,18 @@ def __init__(
if self.text_description:
self.logger.info("Dataset description: %s.", self.text_description)

# Volume mode: return full (T/S, coils, H, W) tensors for dynamic / registration.
self.kspace_context = kspace_context if kspace_context not in (None, 0, False, "") else 0
self.volume_mode = bool(self.kspace_context)
if self.volume_mode and len(spatial_shape) != 3:
raise ValueError(
"FakeMRIBlobsDataset volume mode (kspace_context set) requires "
f"spatial_shape (time_or_slice, height, width). Got {spatial_shape}."
)
self.ndim = 3 if self.volume_mode else 2

self.fake_data: Callable = FakeMRIData(
ndim=len(self.spatial_shape),
ndim=3 if (self.volume_mode or len(self.spatial_shape) == 3) else 2,
blobs_n_samples=kwargs.get("blobs_n_samples", None),
blobs_cluster_std=kwargs.get("blobs_cluster_std", None),
)
Expand All @@ -167,20 +180,18 @@ def __init__(
self.rng = np.random.RandomState()

with temp_seed(self.rng, seed):
# size = sample_size * num_slices if data is 3D
self.data = [
(filename, slice_no, seed)
for (filename, seed) in zip(
self.parse_filenames_data(filenames),
list(self.rng.choice(a=range(int(1e5)), size=self.sample_size, replace=False)),
) # ensure reproducibility
for slice_no in range(self.spatial_shape[0] if len(spatial_shape) == 3 else 1)
]
self.kspace_context = kspace_context if kspace_context else 0
self.ndim = 2 if self.kspace_context == 0 else 3

if self.kspace_context != 0:
raise NotImplementedError("3D reconstruction is not yet supported with FakeMRIBlobsDataset.")
filenames_parsed = self.parse_filenames_data(filenames)
seeds = list(self.rng.choice(a=range(int(1e5)), size=self.sample_size, replace=False))
if self.volume_mode:
# One dataset item per volume.
self.data = list(zip(filenames_parsed, seeds))
else:
# size = sample_size * num_slices if data is 3D (slice-wise 2D)
self.data = [
(filename, slice_no, sample_seed)
for (filename, sample_seed) in zip(filenames_parsed, seeds)
for slice_no in range(self.spatial_shape[0] if len(spatial_shape) == 3 else 1)
]

def parse_filenames_data(self, filenames):
if filenames is None:
Expand All @@ -198,7 +209,10 @@ def parse_filenames_data(self, filenames):
# pylint: disable=logging-fstring-interpolation
self.logger.info(f"Parsing: {(idx + 1) / len(filenames) * 100:.2f}%.")

num_slices = self.spatial_shape[0] if len(self.spatial_shape) == 3 else 1
if self.volume_mode:
num_slices = 1
else:
num_slices = self.spatial_shape[0] if len(self.spatial_shape) == 3 else 1
self.volume_indices[pathlib.PosixPath(filename)] = range(
current_slice_number, current_slice_number + num_slices
)
Expand All @@ -221,7 +235,11 @@ def __len__(self):

def __getitem__(self, index: int) -> dict[str, Any]:
"""Get a sample from the dataset."""
filename, slice_no, sample_seed = self.data[index]
if self.volume_mode:
filename, sample_seed = self.data[index] # ty: ignore[invalid-assignment]
slice_no = 0
else:
filename, slice_no, sample_seed = self.data[index] # ty: ignore[invalid-assignment]

sample = self.fake_data(
sample_size=1,
Expand All @@ -230,7 +248,11 @@ def __getitem__(self, index: int) -> dict[str, Any]:
name=[filename],
seed=sample_seed,
)[0]
sample["kspace"] = sample["kspace"][slice_no]
if self.volume_mode:
# FakeMRIData returns (time/slice, coils, H, W); pipeline expects (coils, time/slice, H, W).
sample["kspace"] = np.swapaxes(sample["kspace"], 0, 1)
else:
sample["kspace"] = sample["kspace"][slice_no]

if "attrs" in sample:
metadata = self._get_metadata(sample["attrs"])
Expand Down
38 changes: 36 additions & 2 deletions direct/data/datasets_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,14 +20,17 @@
from omegaconf import MISSING

from direct.common.subsample_config import MaskingConfig
from direct.config.defaults import BaseConfig
from direct.config import BaseConfig
from direct.data.mri_transforms import (
DemonsFilterType,
HalfSplitType,
MaskSplitterType,
RandomFlipType,
ReconstructionType,
RegistrationSimulateReferenceType,
RescaleMode,
SensitivityMapType,
TransformKey,
TransformsType,
)

Expand Down Expand Up @@ -73,6 +76,25 @@ class NormalizationTransformConfig(BaseConfig):
scale_percentile: float | None = 0.99


@dataclass
class RegistrationTransformConfig(BaseConfig):
registration: bool = False
registration_simulate_reference: RegistrationSimulateReferenceType | None = None
registration_simulate_elastic_sigma: float = 3.0
registration_simulate_elastic_points: int = 3
registration_simulate_elastic_rotate: float = 0.0
registration_simulate_elastic_zoom: float = 0.0
registration_estimate_displacement: bool = True
registration_simulate_reference_from_key_index: int = 0
registration_moving_key: TransformKey = TransformKey.TARGET
demons_filter_type: DemonsFilterType = DemonsFilterType.SYMMETRIC_FORCES
demons_num_iterations: int = 100
demons_smooth_displacement_field: bool = True
demons_standard_deviations: float = 1.5
demons_intensity_difference_threshold: float | None = None
demons_maximum_rms_error: float | None = None


@dataclass
class TransformsConfig(BaseConfig):
"""Configuration for the transforms.
Expand All @@ -81,6 +103,8 @@ class TransformsConfig(BaseConfig):
----------
masking : MaskingConfig
Configuration for the masking.
target_acceleration : float, optional
Target acceleration to override the sampled acceleration with. Default is None.
cropping : CropTransformConfig
Configuration for the cropping.
augmentation : AugmentationTransformConfig
Expand All @@ -95,6 +119,8 @@ class TransformsConfig(BaseConfig):
Configuration for the sensitivity map estimation.
normalization : NormalizationTransformConfig
Configuration for the normalization.
use_acs_as_mask : bool
Use the ACS mask as the sampling mask. Default is False.
delete_acs_mask : bool
Delete ACS mask after its use. Default is True.
delete_kspace : bool
Expand All @@ -107,6 +133,8 @@ class TransformsConfig(BaseConfig):
Default is None.
pad_coils : int, optional
Pad coils. Default is None.
registration : RegistrationTransformConfig
Configuration for the registration transforms.
use_seed : bool
Use seed for the transforms. Typically this should be set to True for reproducibility (e.g. inference),
and False for training. Default is True.
Expand Down Expand Up @@ -136,6 +164,9 @@ class TransformsConfig(BaseConfig):
"""

masking: MaskingConfig | None = field(default_factory=MaskingConfig)
target_acceleration: float | None = None
# Paper adaptive DYNAMIC sampling: independent init/ACS mask per time/slice frame.
dynamic_mask: bool = False
cropping: CropTransformConfig = field(default_factory=CropTransformConfig)
augmentation: AugmentationTransformConfig = field(default_factory=AugmentationTransformConfig)
random_augmentations: RandomAugmentationTransformsConfig = field(default_factory=RandomAugmentationTransformsConfig)
Expand All @@ -145,11 +176,13 @@ class TransformsConfig(BaseConfig):
default_factory=SensitivityMapEstimationTransformConfig
)
normalization: NormalizationTransformConfig = field(default_factory=NormalizationTransformConfig)
use_acs_as_mask: bool = False
delete_acs_mask: bool = True
delete_kspace: bool = True
image_recon_type: ReconstructionType = ReconstructionType.RSS
compress_coils: int | None = None
pad_coils: int | None = None
registration: RegistrationTransformConfig = field(default_factory=RegistrationTransformConfig)
use_seed: bool = True
transforms_type: TransformsType = TransformsType.SUPERVISED
# Next attributes are for the mask splitter in case of transforms_type is set to SSL_SSDU
Expand Down Expand Up @@ -183,7 +216,6 @@ class H5SliceConfig(DatasetConfig):

@dataclass
class CMRxReconConfig(DatasetConfig):
regex_filter: str | None = None
data_root: str | None = None
filenames_filter: list[str] | None = None
filenames_lists: list[str] | None = None
Expand All @@ -207,6 +239,8 @@ class CalgaryCampinasConfig(H5SliceConfig):
@dataclass
class FakeMRIBlobsConfig(DatasetConfig):
pass_attrs: bool = True
# If set (e.g. True / "time"), each sample is a full volume (T/S, coils, H, W).
kspace_context: bool | str | int | None = None


@dataclass
Expand Down
Loading
Loading