Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
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
2 changes: 1 addition & 1 deletion .github/workflows/ci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,7 @@ jobs:
- name: Install dependencies
run: |
uv pip install \
".[bayesian,generative,zarr,dev]" \
".[bayesian,generative,zarr,provenance,dev]" \
monai \
pyro-ppl

Expand Down
2 changes: 1 addition & 1 deletion .github/workflows/guide-notebooks-ec2.yml
Original file line number Diff line number Diff line change
Expand Up @@ -124,7 +124,7 @@ jobs:

# Install nobrainer from checkout on top of the base layer
uv pip install \
".[bayesian,generative,zarr,dev]" \
".[bayesian,generative,zarr,provenance,dev]" \
monai \
pyro-ppl \
matplotlib
Expand Down
3 changes: 3 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -151,3 +151,6 @@ data/
# Model artifacts
*.pth
brain_mask_extraction_model/

# Claude Code review artifacts
.claude/REVIEW.MD/
30 changes: 30 additions & 0 deletions CLAUDE.md
Original file line number Diff line number Diff line change
Expand Up @@ -128,3 +128,33 @@ Every plan MUST include a Constitution Check evaluated before research and after
```

If speckit is not installed, follow the principles above manually.

## Dependency constraint

Do not add external dependencies unless they are already declared in `setup.cfg` or `pyproject.toml`. Current core deps include: torch, monai, nibabel, numpy, click, zarr, fsspec, joblib, scikit-image, psutil. If you think a new dep is justified, flag it explicitly — do not silently add it.

## MONAI-first rule

Before implementing any data loading, transform, loss, metric, or inference utility, check whether MONAI already provides it. Prefer wrapping MONAI over reimplementing. Check `monai.transforms`, `monai.data`, `monai.inferers`, `monai.losses`, `monai.metrics` before writing new code. If MONAI's version is insufficient, document why in a code comment.

## Read before write

Before modifying any existing module:
1. Read the module completely. List its implicit assumptions.
2. Check what MONAI provides that overlaps.
3. Do not guess function signatures — read the source.
4. If the module has existing tests, read those too to understand expected behavior.

## Testing philosophy

Every test should fail if the feature is removed. If a test passes regardless of whether your code exists, it's not testing your code.

## Review process

A reviewer subagent runs automatically (via stop hook) when you finish a task. It executes tests, inspects the diff, and writes a verdict to `.claude/reviews/`. If the reviewer blocks:
1. Read the review file it points to.
2. Fix every CRITICAL finding.
3. Address WARNING findings where reasonable.
4. The hook will re-run on your next stop.

Do not commit until the reviewer passes.
22 changes: 22 additions & 0 deletions conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,28 @@
import torch


def _mps_supports_conv3d() -> bool:
"""Return True if Conv3D works on the MPS backend."""
if not (hasattr(torch.backends, "mps") and torch.backends.mps.is_available()):
return False
try:
m = torch.nn.Conv3d(1, 1, 1).to("mps")
x = torch.randn(1, 1, 2, 2, 2, device="mps")
m(x)
return True
except RuntimeError:
return False


# Disable MPS auto-detection when Conv3D is unsupported (PyTorch < 2.3 on
# Apple Silicon). Without this, ``nobrainer.gpu.get_device()`` returns MPS
# and every 3D convolution raises ``RuntimeError: Conv3D is not supported on
# MPS``.
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
if not _mps_supports_conv3d():
torch.backends.mps.is_available = lambda: False # type: ignore[assignment]


def pytest_collection_modifyitems(config, items):
"""Skip tests marked with @pytest.mark.gpu when CUDA is not available."""
if torch.cuda.is_available():
Expand Down
261 changes: 260 additions & 1 deletion nobrainer/cli/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,12 @@ def cli():
@click.option(
"-v", "--verbose", is_flag=True, help="Print progress messages.", **_option_kwds
)
@click.option(
"--skip-validate",
is_flag=True,
help="Skip preflight input-file validation.",
**_option_kwds,
)
def predict(
*,
infile,
Expand All @@ -117,11 +123,30 @@ def predict(
n_samples,
device,
verbose,
skip_validate,
):
"""Predict labels from a NIfTI volume using a trained PyTorch model.

The predictions are saved to OUTFILE.
"""
# Preflight validation
if not skip_validate:
from ..data.spec import FileStatus, check_file_presence

status = check_file_presence(infile)
if status == FileStatus.NOT_FOUND:
click.echo(click.style(f"ERROR: Input file not found: {infile}", fg="red"))
raise SystemExit(1)
if status == FileStatus.ANNEX_MISSING:
click.echo(
click.style(
f"ERROR: Input is a git-annex symlink with missing content.\n"
f" Run: datalad get {infile}",
fg="red",
)
)
raise SystemExit(1)

if os.path.exists(outfile):
raise FileExistsError(f"Output file already exists: {outfile}")

Expand Down Expand Up @@ -271,7 +296,14 @@ def convert_tfrecords(*, input_paths, output_dir, output_format, verbose):
)
@click.option("--no-conform", is_flag=True, help="Disable auto-conforming.")
@click.option("-v", "--verbose", is_flag=True, help="Print progress.")
def convert_to_zarr(*, output, images, labels, chunk_shape, no_conform, verbose):
@click.option(
"--skip-validate",
is_flag=True,
help="Skip preflight input-file validation.",
)
def convert_to_zarr(
*, output, images, labels, chunk_shape, no_conform, verbose, skip_validate
):
"""Convert NIfTI image+label pairs to a sharded Zarr3 store."""
from ..datasets.zarr_store import create_zarr_store

Expand All @@ -283,6 +315,32 @@ def convert_to_zarr(*, output, images, labels, chunk_shape, no_conform, verbose)
)
sys.exit(1)

# Preflight validation
if not skip_validate:
from ..data.spec import FileStatus, check_file_presence

missing: list[str] = []
annex_missing: list[str] = []
for path in (*images, *labels):
status = check_file_presence(path)
if status == FileStatus.NOT_FOUND:
missing.append(path)
elif status == FileStatus.ANNEX_MISSING:
annex_missing.append(path)
if missing:
click.echo(click.style("ERROR: Files not found:", fg="red"))
for p in missing:
click.echo(f" {p}")
sys.exit(1)
if annex_missing:
click.echo(
click.style("ERROR: git-annex symlinks with missing content:", fg="red")
)
for p in annex_missing:
click.echo(f" {p}")
click.echo(f"\n Run: datalad get {' '.join(annex_missing)}")
sys.exit(1)

pairs = list(zip(images, labels))
chunks = tuple(int(x) for x in chunk_shape.split(","))

Expand Down Expand Up @@ -623,6 +681,126 @@ def info():
click.echo(s)


# ---------------------------------------------------------------------------
# validate / inspect
# ---------------------------------------------------------------------------


@cli.command()
@click.argument("manifest", type=click.Path(exists=True))
@click.option("--json", "as_json", is_flag=True, help="Output results as JSON.")
@click.option("-v", "--verbose", is_flag=True, help="Show per-subject details.")
def validate(*, manifest, as_json, verbose):
"""Validate a dataset manifest against its DataSpec contract.

MANIFEST is the path to a JSON manifest file. Exits non-zero if any
errors are found. Warnings alone do not cause a non-zero exit.
"""
import json as _json

from ..data.spec import DataSpec
from ..data.spec import Severity as _Sev
from ..data.spec import validate as _validate

spec = DataSpec.from_json(manifest)
findings = _validate(spec)

if as_json:
click.echo(_json.dumps([f.to_dict() for f in findings], indent=2))
else:
errs = [f for f in findings if f.severity == _Sev.ERROR]
warns = [f for f in findings if f.severity == _Sev.WARNING]

for f in errs:
click.echo(click.style(f"ERROR {f.field}: {f.message}", fg="red"))
for f in warns:
click.echo(click.style(f"WARN {f.field}: {f.message}", fg="yellow"))

if not findings:
click.echo(click.style("Validation passed.", fg="green"))
else:
click.echo(f"\n{len(errs)} error(s), {len(warns)} warning(s).")

has_errors = any(f.severity == _Sev.ERROR for f in findings)
sys.exit(1 if has_errors else 0)


@cli.command()
@click.argument("path", type=click.Path(exists=True))
@click.option("--json", "as_json", is_flag=True, help="Output as JSON.")
def inspect(*, path, as_json):
"""Inspect NIfTI / Zarr files or a dataset manifest.

PATH may be a JSON manifest, a single NIfTI file, a .zarr store,
or a directory containing NIfTI / Zarr files.
"""
import json as _json
from pathlib import Path as _Path

from ..data.spec import inspect_entry

target = _Path(path)
entries: list[dict] = []

if target.suffix == ".json":
with open(target) as fh:
manifest = _json.load(fh)
base_dir = target.parent
for entry in manifest.get("entries", []):
for key in ("image", "label"):
if key in entry:
p = _Path(entry[key])
if not p.is_absolute():
p = base_dir / p
entries.append(inspect_entry(p))
elif target.is_dir():
for child in sorted(target.iterdir()):
if child.suffix in (".gz", ".nii", ".zarr") or child.suffixes == [
".nii",
".gz",
]:
entries.append(inspect_entry(child))
else:
entries.append(inspect_entry(target))

if as_json:
click.echo(_json.dumps(entries, indent=2))
else:
counts = {"present": 0, "annex_missing": 0, "not_found": 0}
for e in entries:
status = e["file_status"]
counts[status] = counts.get(status, 0) + 1
if status == "present":
shape = tuple(e["shape"]) if e["shape"] else "?"
spacing = tuple(e["spacing"]) if e["spacing"] else "?"
orient = e["orientation"] or "?"
size_mb = (
f"{e['file_size_bytes'] / 1_048_576:.1f}MB"
if e["file_size_bytes"]
else "?"
)
click.echo(
f"{e['path']} PRESENT shape={shape} "
f"spacing={spacing} orient={orient} {size_mb}"
)
elif status == "annex_missing":
click.echo(
click.style(
f"{e['path']} ANNEX_MISSING "
f"(run: datalad get {e['path']})",
fg="yellow",
)
)
else:
click.echo(click.style(f"{e['path']} NOT_FOUND", fg="red"))
click.echo("---")
click.echo(
f"{len(entries)} entries: {counts['present']} present, "
f"{counts['annex_missing']} annex_missing, "
f"{counts['not_found']} not_found"
)


# ---------------------------------------------------------------------------
# zarr subcommands
# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -668,6 +846,87 @@ def zarr_suggest_shards(n_volumes, volume_shape, dtype, n_input_files, levels):
click.echo(json.dumps(result, indent=2))


# ---------------------------------------------------------------------------
# provenance subcommands
# ---------------------------------------------------------------------------


@cli.group()
def provenance():
"""Emit training-run provenance as PROV-O RDF."""


@provenance.command("export")
@click.option(
"--bundle",
required=True,
type=click.Path(exists=True, file_okay=False),
help="Segmentation.save() directory (model.pth + croissant.json).",
**_option_kwds,
)
@click.option(
"--dataspec",
default=None,
type=click.Path(exists=True, dir_okay=False),
help="Optional DataSpec manifest JSON for image+label dataset provenance.",
**_option_kwds,
)
@click.option(
"--out",
default=None,
type=click.Path(),
help="Output file path. Omit to write to stdout.",
)
@click.option(
"--format",
"fmt",
type=click.Choice(["turtle", "json-ld"]),
default="turtle",
help="Serialization format.",
**_option_kwds,
)
@click.option(
"--base-iri",
default="https://neuronets.dev/nobrainer/",
help="Base IRI for minted instance identifiers.",
**_option_kwds,
)
@click.option(
"--agent",
default=None,
help="Optional caller-supplied human/organization agent identity "
"('Name <email>' or an IRI).",
)
@click.option(
"--strict",
is_flag=True,
help="Fail instead of eliding unrepresentable values.",
)
def provenance_export(*, bundle, dataspec, out, fmt, base_iri, agent, strict):
"""Export PROV-O RDF provenance for a saved nobrainer training run."""
from ..provenance import ProvenanceError, export_provenance

try:
text = export_provenance(
bundle,
dataspec_path=dataspec,
fmt=fmt,
base_iri=base_iri,
agent=agent,
strict=strict,
)
except ProvenanceError as exc:
click.echo(click.style(f"ERROR: {exc}", fg="red"))
raise SystemExit(1) from exc

if out:
with open(out, "w") as fh:
fh.write(text)
click.echo(f"Wrote {fmt} provenance to {out}")
else:
click.echo(text)


# For debugging only.
if __name__ == "__main__":
cli()
1 change: 1 addition & 0 deletions nobrainer/data/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
"""Dataset specifications and static data for nobrainer."""
Loading
Loading