Skip to content
Merged
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
174 changes: 79 additions & 95 deletions deepcell_types/utils/__init__.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,15 @@
"""Utilities for model / data access.

The registries below pin checksums of the paper-release checkpoints;
uploads to ``users.deepcell.org`` use the asset paths constructed below
(``models/<filename>``). The hash algorithm is auto-detected from the
digest length (32 hex → md5, 64 hex → sha256), so entries can be migrated
to the stronger sha256 individually; new entries should pin sha256. Some
baselines ship more than one file (e.g. ``maps`` needs its ``_stats.npz``
companion; ``xgboost`` needs its ``.remap.json`` label-remap), so each
baseline entry is a list of ``(filename, hash)`` tuples.
Model, baseline, and training-data downloads are delegated to the shared
``deepcell-auth`` client, whose bundled ``asset_manifest.yaml`` is the single
source of truth for asset keys and integrity hashes. The functions here are
thin adapters that keep this package's public API: the ``ValueError`` on an
unknown identifier, the pointer to Nimbus' upstream weights, and the
``Path`` / ``list[Path]`` return types.

The version and baseline names below are mirrored statically from that
manifest so an unknown identifier is rejected without a manifest read;
``tests/test_download_delegation.py`` fails if they drift.
"""

__all__ = [
Expand All @@ -21,70 +23,33 @@
"resolve_supported_marker",
]

_latest = "2026-06-15"
_DEFAULT_MODEL_VERSION = "2026-06-15"

# Main model checkpoints. Values are ``(asset_filename, md5)``.
# Two released residual-MLP checkpoints share the current vocabulary (51 cell
# types, 278 markers) and load via stock ``predict.py`` (the resMLP head is
# Released residual-MLP checkpoints, sharing the current vocabulary (51 cell
# types, 278 markers) and loading via stock ``predict.py`` (the resMLP head is
# auto-detected):
# - ``2026-06-15`` (default): the from-scratch resMLP (paper Fig-3c headline).
# - ``2026-06-23``: the pretrain -> finetune (SSL) resMLP arm.
#
# Each md5 tracks the vocab-bundled asset (carries ``ct2idx`` +
# ``canonical_channels`` so ``validate_checkpoint_vocabulary`` can verify
# ordering). The original un-bundled assets predated the guard and raised
# ``ValueError: ... does not bundle a ct2idx`` on every ``predict()`` call;
# they were re-packaged with ``scripts/repackage_release_checkpoint.py``. Each
# bundled asset must be uploaded at its filename below before its md5 is served.
_model_registry = {
"2026-06-15": (
"deepcell-types_2026-06-15_resmlp.pt",
"b819a7e0b177ad5330394eab3c6c7ad8",
),
"2026-06-23": (
"deepcell-types_2026-06-23_resmlp_ptft.pt",
"402e94c103c5e489a57433cd107009d3",
),
}

# Baseline-model checkpoints. Values are lists of ``(asset_filename, md5)``.
# Single-file baselines have a one-element list; ``maps`` and ``xgboost``
# additionally ship a companion file required at inference. The md5s track the
# 2026-06-30 DCT-sampler retrains -- the release-of-record that produced the
# published baseline prediction CSVs and headline numbers.
# - ``2026-06-23-ptft``: the pretrain -> finetune (SSL) resMLP arm.
_MODEL_VERSIONS = ("2026-06-15", "2026-06-23-ptft")

# Also in the manifest, but deliberately not offered here: the 2025 CLIP
# checkpoints predate the current architecture and cannot be loaded by this
# code. They stay served for reproducibility with the matching historical
# commit, so ``download_model`` rejects them rather than handing back a
# checkpoint that fails at ``load_state_dict``.
_LEGACY_MODEL_VERSIONS = ("2025-06-09", "2025-06-09_public-data-only")

# All three baselines are served as a single ``.tar.gz`` bundle, one
# subdirectory per baseline, holding each one's weights plus any companion file
# required at inference (maps -> ``_stats.npz``; xgboost -> ``.remap.json``).
# Requesting any one baseline therefore downloads the whole bundle; the other
# two are then served from the cache without a second transfer.
#
# Nimbus is intentionally absent: it is inference-only with pretrained weights
# produced and distributed upstream (angelolab/Nimbus-Inference), which this
# project does not redistribute. ``download_baseline_checkpoint("nimbus")``
# raises with a pointer to the official source instead (see below).
_baseline_registry = {
"cellsighter": [
(
"deepcell-types_baseline-cellsighter.pth",
"c5a105fb044ad82c82817a32aabbae7c",
),
],
"maps": [
(
"deepcell-types_baseline-maps.pth",
"1ad5e30ceaccfdb91050b663099258fb",
),
(
"deepcell-types_baseline-maps_stats.npz",
"1f462d0c8bf531af73026d415d52728d",
),
],
"xgboost": [
(
"deepcell-types_baseline-xgboost.json",
"9e51cf1af8c6c43b00871a821fde0f57",
),
(
"deepcell-types_baseline-xgboost.remap.json",
"fba1b0e705e5f7747eb2f9cae30815ba",
),
],
}
_BASELINE_NAMES = ("cellsighter", "maps", "xgboost")

# Nimbus pretrained weights are distributed upstream (via the
# ``nimbus-inference`` library / Hugging Face Hub), not re-hosted here.
Expand All @@ -100,38 +65,39 @@ def download_model(*, version=None):
----------
version : str, optional
Which checkpoint version to download. Defaults to ``None``,
which resolves to the most-recently-released version
(``_latest`` in this module).
which resolves to the default version
(``_DEFAULT_MODEL_VERSION`` in this module).

Returns
-------
pathlib.Path
Local path to the downloaded checkpoint.
"""
from ._auth import fetch_data
from deepcell_auth import download_deepcell_types_model

version = version if version is not None else _latest
if version not in _model_registry:
version = version if version is not None else _DEFAULT_MODEL_VERSION
if version not in _MODEL_VERSIONS:
raise ValueError(
f"Unknown model version {version!r}. "
f"Known versions: {sorted(_model_registry)}."
f"Known versions: {sorted(_MODEL_VERSIONS)}."
)
filename, md5 = _model_registry[version]
return fetch_data(f"models/{filename}", cache_subdir="models", file_hash=md5)
return download_deepcell_types_model(version)


def list_model_versions():
"""Return the available pre-trained model versions, newest first.
"""Return the available pre-trained model versions, default first.

Returns
-------
list of str
Version identifiers accepted by :func:`download_model`. The first
element is always the default (``_latest``) version that
``download_model()`` resolves to with no argument.
element is always the default version that ``download_model()``
resolves to with no argument.
"""
others = sorted((v for v in _model_registry if v != _latest), reverse=True)
return [_latest, *others]
others = sorted(
(v for v in _MODEL_VERSIONS if v != _DEFAULT_MODEL_VERSION), reverse=True
)
return [_DEFAULT_MODEL_VERSION, *others]


def download_baseline_checkpoint(name):
Expand All @@ -155,13 +121,15 @@ def download_baseline_checkpoint(name):
Returns
-------
list[pathlib.Path]
Local paths to every file downloaded for this baseline, in the
order declared in ``_baseline_registry``. Note the asymmetry with
:func:`download_model`, which returns a single ``Path``: baselines
return a *list* because some ship companion files. Call
:func:`list_baseline_names` for the accepted identifiers.
Local paths to every checkpoint file for this baseline, sorted by
filename. Note the asymmetry with :func:`download_model`, which
returns a single ``Path``: baselines return a *list* because some
ship companion files. Call :func:`list_baseline_names` for the
accepted identifiers.
"""
from ._auth import fetch_data
from deepcell_auth import download_deepcell_types_baseline

from ._archive import extract_archive

if name == "nimbus":
raise ValueError(
Expand All @@ -171,14 +139,31 @@ def download_baseline_checkpoint(name):
"Python 3.11), which downloads the weights automatically; see "
f"{_NIMBUS_UPSTREAM_URL}."
)
if name not in _baseline_registry:
if name not in _BASELINE_NAMES:
raise ValueError(
f"Unknown baseline {name!r}. Known baselines: {sorted(_BASELINE_NAMES)}."
)

(archive,) = download_deepcell_types_baseline(name)
bundle_dir = archive.parent / archive.name.removesuffix(".tar.gz")
# Mirrors cellSAM: the bundle is unpacked once, and its presence is what
# skips re-extraction on later calls. ``fetch_data`` still re-checks the
# archive's pinned hash on every call, so a corrupt *download* is caught;
# a hand-edited extraction directory is not.
if not bundle_dir.is_dir():
extract_archive(archive, archive.parent)
baseline_dir = bundle_dir / name
contents = (
sorted(p for p in baseline_dir.iterdir() if p.is_file())
if baseline_dir.is_dir()
else []
)
if not contents:
raise ValueError(
f"Unknown baseline {name!r}. Known baselines: {sorted(_baseline_registry)}."
f"Baseline bundle {archive.name} did not unpack to a directory of "
f"checkpoint files at {baseline_dir}. Delete {bundle_dir} and retry."
)
return [
fetch_data(f"models/{filename}", cache_subdir="models", file_hash=md5)
for filename, md5 in _baseline_registry[name]
]
return contents


def list_baseline_names():
Expand All @@ -190,7 +175,7 @@ def list_baseline_names():
Names accepted by :func:`download_baseline_checkpoint` (the
``list``-returning counterpart to :func:`list_model_versions`).
"""
return sorted(_baseline_registry)
return sorted(_BASELINE_NAMES)


def list_supported_markers(*, zarr_path=None):
Expand Down Expand Up @@ -269,9 +254,6 @@ def list_supported_cell_types(*, zarr_path=None):
return sorted(DCTConfig(zarr_path=zarr_path).ct2idx)


_training_data_asset_key = "data/deepcell-types/public_data_v1.1.zip"


def download_training_data(*, extract=False):
"""Download the public training-data corpus for deepcell-types (v1.1).

Expand All @@ -297,9 +279,11 @@ def download_training_data(*, extract=False):
Local path to the downloaded ``.zip`` (or, when ``extract=True``,
the directory it was extracted into).
"""
from ._auth import extract_archive, fetch_data
from deepcell_auth import download_deepcell_types_data

from ._archive import extract_archive

zip_path = fetch_data(_training_data_asset_key, cache_subdir="data")
zip_path = download_deepcell_types_data()
if extract:
return extract_archive(zip_path)
return zip_path
120 changes: 120 additions & 0 deletions deepcell_types/utils/_archive.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,120 @@
"""Path-traversal-safe extraction of downloaded archives.

Kept in this package rather than delegated to ``deepcell_auth``, whose
``extract_archive`` calls ``extractall`` without vetting members.
"""

import logging
import tarfile
import zipfile
from pathlib import Path

logger = logging.getLogger(__name__)

# Deliberately generous limits accommodate the public multi-terabyte corpus
# while still bounding hostile or accidentally unbounded archives.
MAX_ARCHIVE_MEMBERS = 100_000
MAX_ARCHIVE_MEMBER_BYTES = 2 * 2**40
MAX_ARCHIVE_TOTAL_BYTES = 10 * 2**40


def extract_archive(
archive_path,
dest=None,
*,
max_members=MAX_ARCHIVE_MEMBERS,
max_member_bytes=MAX_ARCHIVE_MEMBER_BYTES,
max_total_bytes=MAX_ARCHIVE_TOTAL_BYTES,
):
"""Safely extract a ``.zip`` or ``.tar[.gz]`` archive.

Rejects members whose resolved path would escape ``dest`` (zip-slip /
tar path-traversal) and rejects tar symlink/hardlink members, so it is
safe to call on archives downloaded from a remote source.

Parameters
----------
archive_path : str or pathlib.Path
Path to a ``.zip`` or ``.tar``/``.tar.gz`` archive.
dest : str or pathlib.Path, optional
Destination directory. Defaults to the archive's parent directory.

Returns
-------
pathlib.Path
The destination directory the archive was extracted into.
"""
archive_path = Path(archive_path)
dest = Path(dest) if dest is not None else archive_path.parent
dest.mkdir(parents=True, exist_ok=True)
dest_resolved = dest.resolve()

def _within(name):
try:
(dest / name).resolve().relative_to(dest_resolved)
return True
except ValueError:
return False

if zipfile.is_zipfile(archive_path):
with zipfile.ZipFile(archive_path) as zf:
infos = zf.infolist()
if len(infos) > max_members:
raise ValueError(
f"Refusing to extract archive with {len(infos)} members; "
f"limit is {max_members}."
)
total = 0
for member in infos:
is_symlink = (member.external_attr >> 16) & 0o170000 == 0o120000
if not _within(member.filename):
raise ValueError(
f"Refusing to extract {member.filename!r}: path escapes {dest}."
)
if is_symlink:
raise ValueError(
f"Refusing to extract unsafe zip member {member.filename!r}."
)
if member.file_size > max_member_bytes:
raise ValueError(
f"Refusing to extract {member.filename!r}: declared size "
f"exceeds {max_member_bytes} bytes."
)
total += member.file_size
if total > max_total_bytes:
raise ValueError(
f"Refusing to extract archive larger than "
f"{max_total_bytes} bytes."
)
zf.extractall(dest)
elif tarfile.is_tarfile(archive_path):
with tarfile.open(archive_path) as tf:
members = tf.getmembers()
if len(members) > max_members:
raise ValueError(
f"Refusing to extract archive with {len(members)} members; "
f"limit is {max_members}."
)
total = 0
for member in members:
if not (member.isfile() or member.isdir()) or not _within(member.name):
raise ValueError(
f"Refusing to extract unsafe tar member {member.name!r}."
)
if member.size > max_member_bytes:
raise ValueError(
f"Refusing to extract {member.name!r}: declared size "
f"exceeds {max_member_bytes} bytes."
)
total += member.size
if total > max_total_bytes:
raise ValueError(
f"Refusing to extract archive larger than "
f"{max_total_bytes} bytes."
)
tf.extractall(dest, filter="data")
else:
raise ValueError(f"{archive_path} is not a recognized .zip or .tar archive.")

logger.info(f"Extracted {archive_path} to {dest}")
return dest
Loading
Loading