diff --git a/.gitignore b/.gitignore
index 9a00cf8..7b10493 100644
--- a/.gitignore
+++ b/.gitignore
@@ -233,8 +233,10 @@ src/modernmolbert/eval/junk/
.prompt.md
.agents/
/*.prompt.md
-repomix*.xml
+/*plan*.md
+*repomix*.xml
CLAUDE.md
+.claude/
# eval
tmp_eval/
diff --git a/src/modernmolbert/eval/cache.py b/src/modernmolbert/eval/cache.py
index fc40c40..e202421 100644
--- a/src/modernmolbert/eval/cache.py
+++ b/src/modernmolbert/eval/cache.py
@@ -71,11 +71,7 @@ def _public_featurizer_params(featurizer: RepresentationFeaturizer) -> dict[str,
if is_dataclass(featurizer):
params = asdict(featurizer)
else:
- params = {
- key: value
- for key, value in vars(featurizer).items()
- if not key.startswith("_")
- }
+ params = {key: value for key, value in vars(featurizer).items() if not key.startswith("_")}
# Exclude heavy/runtime objects if present.
for key in [
@@ -96,9 +92,7 @@ def featurizer_cache_identity(
return {
"name": featurizer.name,
- "class": (
- f"{featurizer.__class__.__module__}.{featurizer.__class__.__qualname__}"
- ),
+ "class": (f"{featurizer.__class__.__module__}.{featurizer.__class__.__qualname__}"),
"params": _public_featurizer_params(featurizer),
}
@@ -227,8 +221,7 @@ def _validate_cache_metadata(
observed = metadata.get(key)
if observed != expected:
raise ValueError(
- f"Cache metadata mismatch for {key!r}: "
- f"expected {expected!r}, observed {observed!r}"
+ f"Cache metadata mismatch for {key!r}: expected {expected!r}, observed {observed!r}"
)
@@ -251,9 +244,7 @@ def get_or_compute_features(
"""
if smiles_column not in frame.columns:
- raise ValueError(
- f"Split {split_name!r} is missing SMILES column {smiles_column!r}"
- )
+ raise ValueError(f"Split {split_name!r} is missing SMILES column {smiles_column!r}")
smiles_values = frame[smiles_column].tolist()
molecule_hash = hash_molecule_values(smiles_values)
diff --git a/src/modernmolbert/eval/cli/prepare_moleculenet.py b/src/modernmolbert/eval/cli/prepare_moleculenet.py
index 1564c11..e8beae1 100644
--- a/src/modernmolbert/eval/cli/prepare_moleculenet.py
+++ b/src/modernmolbert/eval/cli/prepare_moleculenet.py
@@ -8,8 +8,7 @@
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description=(
- "Prepare DeepChem/MoleculeNet datasets as local sanitized "
- "SMILES/SELFIES Parquet files."
+ "Prepare DeepChem/MoleculeNet datasets as local sanitized SMILES/SELFIES Parquet files."
)
)
@@ -18,10 +17,7 @@ def parse_args() -> argparse.Namespace:
nargs="+",
default=list(CORE_SPECS),
choices=sorted(ALL_SPECS),
- help=(
- "Dataset names to prepare. Defaults to the core suite: "
- + ", ".join(CORE_SPECS)
- ),
+ help=("Dataset names to prepare. Defaults to the core suite: " + ", ".join(CORE_SPECS)),
)
parser.add_argument(
"--output_root",
diff --git a/src/modernmolbert/eval/moleculenet.py b/src/modernmolbert/eval/moleculenet.py
index 4dc0a3d..936811b 100644
--- a/src/modernmolbert/eval/moleculenet.py
+++ b/src/modernmolbert/eval/moleculenet.py
@@ -126,9 +126,7 @@ def prepare_many(
for dataset_name in dataset_names:
if dataset_name not in ALL_SPECS:
valid = ", ".join(sorted(ALL_SPECS))
- raise ValueError(
- f"Unknown dataset {dataset_name!r}. Valid choices: {valid}"
- )
+ raise ValueError(f"Unknown dataset {dataset_name!r}. Valid choices: {valid}")
prepare_dataset(
spec=ALL_SPECS[dataset_name],
@@ -294,12 +292,9 @@ def prepare_dataset(
"valid": frac_valid,
"test": frac_test,
},
- "scaffold_stats": compute_scaffold_stats(split_frame)
- if split == "scaffold"
- else None,
+ "scaffold_stats": compute_scaffold_stats(split_frame) if split == "scaffold" else None,
"split_scaffold_stats": {
- split_name: compute_scaffold_stats(split_df)
- for split_name, split_df in splits.items()
+ split_name: compute_scaffold_stats(split_df) for split_name, split_df in splits.items()
}
if split == "scaffold"
else None,
@@ -374,9 +369,7 @@ def deepchem_dataset_to_frame(dataset: Any, tasks: Sequence[str]) -> pd.DataFram
if y.shape[1] == 1 and len(tasks) == 0:
tasks = ["label"]
else:
- raise ValueError(
- f"Task count mismatch: len(tasks)={len(tasks)}, y.shape={y.shape}"
- )
+ raise ValueError(f"Task count mismatch: len(tasks)={len(tasks)}, y.shape={y.shape}")
rows: list[dict[str, Any]] = []
@@ -580,9 +573,7 @@ def split_sanitized_frame(
frac_valid=frac_valid,
)
- raise ValueError(
- f"Unsupported local split {split!r}. Use 'scaffold', 'random', or 'index'."
- )
+ raise ValueError(f"Unsupported local split {split!r}. Use 'scaffold', 'random', or 'index'.")
def random_split_frame(
@@ -811,9 +802,7 @@ def compute_duplicate_stats(frame: pd.DataFrame) -> dict[str, Any]:
n_duplicate_rows = int(n_valid - n_unique)
duplicated_values = valid[valid.duplicated(keep=False)]
- duplicate_group_sizes = (
- duplicated_values.value_counts().sort_values(ascending=False).tolist()
- )
+ duplicate_group_sizes = duplicated_values.value_counts().sort_values(ascending=False).tolist()
return {
"n_valid_rows": n_valid,
@@ -836,8 +825,7 @@ def grouped_random_split_frame(
rng = np.random.default_rng(seed)
groups = [
- indices.to_list()
- for _, indices in frame.groupby(group_column, sort=False).groups.items()
+ indices.to_list() for _, indices in frame.groupby(group_column, sort=False).groups.items()
]
rng.shuffle(groups)
diff --git a/src/modernmolbert/paths.py b/src/modernmolbert/paths.py
index f4f44ad..617a709 100644
--- a/src/modernmolbert/paths.py
+++ b/src/modernmolbert/paths.py
@@ -19,9 +19,7 @@ def find_project_root(
if env_root:
root = Path(env_root).expanduser().resolve()
if not root.exists():
- raise FileNotFoundError(
- f"MODERNMOLBERT_ROOT points to a missing path: {root}"
- )
+ raise FileNotFoundError(f"MODERNMOLBERT_ROOT points to a missing path: {root}")
return root
if start is None:
@@ -44,9 +42,7 @@ def find_project_root(
)
-def project_path(
- *parts: str | os.PathLike[str], start: str | Path | None = None
-) -> Path:
+def project_path(*parts: str | os.PathLike[str], start: str | Path | None = None) -> Path:
"""Return an absolute path inside the project root."""
return find_project_root(start=start).joinpath(*parts)
@@ -58,9 +54,7 @@ def data_path(*parts: str | os.PathLike[str], start: str | Path | None = None) -
return project_path("data", *parts, start=start)
-def outputs_path(
- *parts: str | os.PathLike[str], start: str | Path | None = None
-) -> Path:
+def outputs_path(*parts: str | os.PathLike[str], start: str | Path | None = None) -> Path:
"""Return an absolute path inside the project outputs directory."""
return project_path("outputs", *parts, start=start)
diff --git a/src/modernmolbert/train_selfies_ape_modernbert.py b/src/modernmolbert/train_selfies_ape_modernbert.py
index ae23b34..159d3f7 100644
--- a/src/modernmolbert/train_selfies_ape_modernbert.py
+++ b/src/modernmolbert/train_selfies_ape_modernbert.py
@@ -9,14 +9,15 @@
import argparse
import hashlib
+import re
import time
import json
import math
import platform
import random
-from dataclasses import dataclass
+from dataclasses import dataclass, field
from pathlib import Path
-from typing import Any
+from typing import Any, ClassVar
import os
from dotenv import load_dotenv
@@ -153,7 +154,61 @@ def parse_args() -> argparse.Namespace:
)
# MLM
- parser.add_argument("--mlm_probability", type=float, default=0.30)
+ parser.add_argument(
+ "--mlm_probability",
+ type=float,
+ default=0.30,
+ help=(
+ "Fraction of eligible tokens to mask. For span/hetero_span strategies the "
+ "budget is round(n_eligible × mlm_probability); short sequences may exceed "
+ "this rate when a single span covers the full budget in one draw."
+ ),
+ )
+ parser.add_argument(
+ "--masking_strategy",
+ type=str,
+ choices=["standard", "span", "hetero_span"],
+ default="standard",
+ help=(
+ "MLM masking strategy. "
+ "'standard': independent Bernoulli per token (original). "
+ "'span': budget-based contiguous APE-token span masking. "
+ "'hetero_span': span masking with span-start positions weighted toward "
+ "APE tokens that contain heteroatoms (N, O, S, P, F, Cl, Br, I, Se, Si)."
+ ),
+ )
+ parser.add_argument(
+ "--span_p_geom",
+ type=float,
+ default=0.4,
+ help=(
+ "Success probability for the geometric distribution used to sample span lengths "
+ "(mean span length = 1/span_p_geom ≈ 2.5 APE tokens at default 0.4). "
+ "Only used when --masking_strategy is 'span' or 'hetero_span'."
+ ),
+ )
+ parser.add_argument(
+ "--span_max_length",
+ type=int,
+ default=6,
+ help=(
+ "Maximum span length in APE tokens. Individual sampled lengths are clamped to "
+ "this value. Adjacent independent spans can form longer contiguous masked runs — "
+ "this parameter bounds individual draws, not total run length. "
+ "Only used when --masking_strategy is 'span' or 'hetero_span'."
+ ),
+ )
+ parser.add_argument(
+ "--heteroatom_start_weight",
+ type=float,
+ default=2.0,
+ help=(
+ "Sampling weight multiplier for span-start positions whose APE token contains "
+ "a heteroatom bracket (N, O, S, P, F, Cl, Br, I, Se, Si). "
+ "Non-heteroatom-containing positions receive weight 1.0. "
+ "Only used when --masking_strategy is 'hetero_span'."
+ ),
+ )
# Training
parser.add_argument("--max_steps", type=int, default=150_000)
@@ -278,6 +333,15 @@ def validate_args(args: argparse.Namespace, backend: str) -> None:
"--load_best_model_at_end requires --save_steps to equal --eval_steps "
"so every evaluated checkpoint can be selected as best."
)
+ if args.masking_strategy not in {"standard", "span", "hetero_span"}:
+ raise ValueError(f"Unknown masking_strategy: {args.masking_strategy!r}")
+ if args.masking_strategy in {"span", "hetero_span"}:
+ if not (0.0 < args.span_p_geom < 1.0):
+ raise ValueError("span_p_geom must be in (0, 1)")
+ if args.span_max_length < 1:
+ raise ValueError("span_max_length must be >= 1")
+ if args.masking_strategy == "hetero_span" and args.heteroatom_start_weight <= 0.0:
+ raise ValueError("heteroatom_start_weight must be > 0")
def adjust_args_for_backend(args: argparse.Namespace, backend: str) -> argparse.Namespace:
@@ -617,13 +681,45 @@ class MolecularMLMCollator:
vocab_size: int
mlm_probability: float
special_token_ids: list[int]
+ masking_strategy: str = "standard"
+ span_p_geom: float = 0.4
+ span_max_length: int = 6
+ heteroatom_start_weight: float = 2.0
+ ids_to_tokens: dict[int, str] = field(default_factory=dict)
+
+ # ClassVar: excluded from __init__ by dataclass machinery.
+ # Ordered longest-first so alternation matches Cl before C, Br before B, Se before S.
+ _HETEROATOM_IN_BRACKET: ClassVar[re.Pattern] = re.compile(
+ r"\["
+ r"[=#/\\@+\-]*"
+ r"(?:Cl|Br|Se|Si|[NOSPFI])"
+ r"[^\]]*"
+ r"\]"
+ )
def __post_init__(self) -> None:
special_ids = {int(token_id) for token_id in self.special_token_ids}
eligible = [token_id for token_id in range(self.vocab_size) if token_id not in special_ids]
-
self._eligible_replacement_ids = torch.tensor(eligible, dtype=torch.long)
+ if self.masking_strategy in {"span", "hetero_span"}:
+ if not (0.0 < self.span_p_geom < 1.0):
+ raise ValueError("span_p_geom must be in (0, 1)")
+ if self.span_max_length < 1:
+ raise ValueError("span_max_length must be >= 1")
+ self._geom_dist: torch.distributions.Geometric | None = torch.distributions.Geometric(
+ torch.tensor(self.span_p_geom)
+ )
+ else:
+ self._geom_dist = None
+
+ if self.masking_strategy == "hetero_span":
+ if self.heteroatom_start_weight <= 0.0:
+ raise ValueError("heteroatom_start_weight must be > 0")
+ self._token_start_weights = self._build_token_start_weights()
+ else:
+ self._token_start_weights = None
+
def __call__(self, examples: list[dict[str, Any]]) -> dict[str, torch.Tensor]:
ids = [
torch.tensor(ex["input_ids"], dtype=torch.long)
@@ -638,25 +734,16 @@ def __call__(self, examples: list[dict[str, Any]]) -> dict[str, torch.Tensor]:
attention_mask = (input_ids != self.pad_token_id).long()
labels = input_ids.clone()
- probability_matrix = torch.full(labels.shape, self.mlm_probability)
-
special_mask = torch.zeros_like(labels, dtype=torch.bool)
for sid in self.special_token_ids:
special_mask |= labels.eq(sid)
- probability_matrix.masked_fill_(special_mask, 0.0)
- probability_matrix.masked_fill_(attention_mask.eq(0), 0.0)
-
- masked_indices = torch.bernoulli(probability_matrix).bool()
-
- if self.mlm_probability > 0.0 and not masked_indices.any():
- eligible_positions = (~special_mask & attention_mask.bool()).nonzero(as_tuple=False)
- if len(eligible_positions) > 0:
- idx = int(torch.randint(len(eligible_positions), (1,)).item())
- row_pos = eligible_positions[idx]
- row = int(row_pos[0].item())
- col = int(row_pos[1].item())
- masked_indices[row, col] = True
+ if self.masking_strategy == "standard":
+ masked_indices = self._sample_standard_mask(labels, attention_mask, special_mask)
+ elif self.masking_strategy in {"span", "hetero_span"}:
+ masked_indices = self._sample_batch_span_mask(input_ids, attention_mask, special_mask)
+ else:
+ raise ValueError(f"Unknown masking_strategy: {self.masking_strategy!r}")
labels[~masked_indices] = -100
@@ -688,6 +775,51 @@ def __call__(self, examples: list[dict[str, Any]]) -> dict[str, torch.Tensor]:
"labels": labels,
}
+ def _sample_standard_mask(
+ self,
+ labels: torch.Tensor,
+ attention_mask: torch.Tensor,
+ special_mask: torch.Tensor,
+ ) -> torch.Tensor:
+ probability_matrix = torch.full(labels.shape, self.mlm_probability)
+ probability_matrix.masked_fill_(special_mask, 0.0)
+ probability_matrix.masked_fill_(attention_mask.eq(0), 0.0)
+ masked_indices = torch.bernoulli(probability_matrix).bool()
+
+ if self.mlm_probability > 0.0 and not masked_indices.any():
+ eligible_positions = (~special_mask & attention_mask.bool()).nonzero(as_tuple=False)
+ if len(eligible_positions) > 0:
+ idx = int(torch.randint(len(eligible_positions), (1,)).item())
+ row_pos = eligible_positions[idx]
+ row = int(row_pos[0].item())
+ col = int(row_pos[1].item())
+ masked_indices[row, col] = True
+
+ return masked_indices
+
+ def _sample_batch_span_mask(
+ self,
+ input_ids: torch.Tensor,
+ attention_mask: torch.Tensor,
+ special_mask: torch.Tensor,
+ ) -> torch.Tensor:
+ batch_size = input_ids.size(0)
+ masked_indices = torch.zeros_like(input_ids, dtype=torch.bool)
+ for i in range(batch_size):
+ row_mask = self._sample_span_mask(
+ input_ids_row=input_ids[i],
+ attention_mask_row=attention_mask[i],
+ special_mask_row=special_mask[i],
+ )
+ if not row_mask.any():
+ eligible = (~special_mask[i] & attention_mask[i].bool()).nonzero(as_tuple=False)
+ if len(eligible) > 0:
+ rand_idx = int(torch.randint(len(eligible), (1,)).item())
+ col = int(eligible[rand_idx].item())
+ row_mask[col] = True
+ masked_indices[i] = row_mask
+ return masked_indices
+
def eligible_random_token_ids(
self,
device: torch.device | None = None,
@@ -703,6 +835,105 @@ def eligible_random_token_ids(
return self._eligible_replacement_ids.to(device)
+ def _build_token_start_weights(self) -> torch.Tensor:
+ """Return a (vocab_size,) float weight tensor for heteroatom-biased span starts.
+
+ Covered heteroatom set: N, O, S, P, F, Cl, Br, I, Se, Si.
+ Elements not in this set (e.g. B, Sn, As, Ge) receive weight 1.0.
+ If intentional coverage of additional elements is needed, extend
+ _HETEROATOM_IN_BRACKET accordingly.
+
+ Token IDs in special_token_ids receive weight 0.0 as a defensive guard;
+ the eligible-position filter in _sample_span_mask is the primary barrier.
+ Tokens matching the heteroatom pattern receive weight heteroatom_start_weight.
+ All other tokens receive weight 1.0.
+ """
+ weights = torch.ones(self.vocab_size, dtype=torch.float32)
+ special_ids = set(self.special_token_ids)
+ for tok_id, tok_str in self.ids_to_tokens.items():
+ if tok_id in special_ids:
+ weights[tok_id] = 0.0
+ elif self._HETEROATOM_IN_BRACKET.search(tok_str):
+ weights[tok_id] = float(self.heteroatom_start_weight)
+ return weights
+
+ def _sample_span_mask(
+ self,
+ input_ids_row: torch.Tensor,
+ attention_mask_row: torch.Tensor,
+ special_mask_row: torch.Tensor,
+ ) -> torch.Tensor:
+ """Sample a span-based boolean mask for one sequence.
+
+ Contiguous spans of APE tokens are sampled until the number of newly
+ masked positions reaches round(n_eligible × mlm_probability).
+ Span lengths are drawn from a Geometric(span_p_geom) distribution and
+ clamped to span_max_length. For hetero_span, span-start positions are
+ sampled with weights proportional to heteroatom content.
+
+ Adjacent independent spans may produce contiguous masked runs longer than
+ span_max_length — the parameter bounds individual draws, not total runs.
+ On very short sequences the actual masked fraction may exceed mlm_probability
+ because a single span can cover the entire budget in one draw.
+ """
+ seq_len = input_ids_row.size(0)
+ masked = torch.zeros(seq_len, dtype=torch.bool)
+
+ eligible_mask = (~special_mask_row) & attention_mask_row.bool()
+ eligible_pos = eligible_mask.nonzero(as_tuple=False).squeeze(1)
+
+ if len(eligible_pos) == 0:
+ return masked
+
+ n_eligible = len(eligible_pos)
+ budget = max(1, round(n_eligible * self.mlm_probability))
+
+ # Pre-sample all geometric span lengths in one vectorised call.
+ max_draws = budget * 5
+ assert self._geom_dist is not None
+ span_lengths = (
+ self._geom_dist.sample((max_draws,)).long() + 1 # shift k≥0 → k≥1
+ ).clamp(max=self.span_max_length)
+
+ if self.masking_strategy == "hetero_span" and self._token_start_weights is not None:
+ tok_ids_at_eligible = input_ids_row[eligible_pos]
+ pos_weights = self._token_start_weights[tok_ids_at_eligible].clone()
+ else:
+ pos_weights = torch.ones(n_eligible, dtype=torch.float32)
+
+ masked_count = 0
+
+ for draw_idx in range(max_draws):
+ if masked_count >= budget:
+ break
+ if pos_weights.sum().item() == 0.0:
+ break
+
+ start_local = int(torch.multinomial(pos_weights, num_samples=1).item())
+ start = int(eligible_pos[start_local].item())
+ span_len = int(span_lengths[draw_idx].item())
+ end = min(start + span_len, seq_len)
+
+ for pos in range(start, end):
+ if not attention_mask_row[pos].item() or special_mask_row[pos].item():
+ end = pos
+ break
+
+ if end <= start:
+ pos_weights[start_local] = 0.0
+ continue
+
+ new_count = int((~masked[start:end]).sum().item())
+ masked[start:end] = True
+ masked_count += new_count
+
+ # Zero weights for covered eligible positions so subsequent draws
+ # explore unmasked territory.
+ covered = (eligible_pos >= start) & (eligible_pos < end)
+ pos_weights[covered] = 0.0
+
+ return masked
+
MODERNBERT_CONFIGS = {
"base": "answerdotai/ModernBERT-base",
@@ -792,6 +1023,12 @@ def log_training_plan(
print(f" max_steps: {args.max_steps}", flush=True)
print(f" max_seq_length: {args.max_seq_length}", flush=True)
print(f" mlm_probability: {args.mlm_probability}", flush=True)
+ print(f" masking_strategy: {args.masking_strategy}", flush=True)
+ if args.masking_strategy in {"span", "hetero_span"}:
+ print(f" span_p_geom: {args.span_p_geom}", flush=True)
+ print(f" span_max_length: {args.span_max_length}", flush=True)
+ if args.masking_strategy == "hetero_span":
+ print(f" heteroatom_start_weight: {args.heteroatom_start_weight}", flush=True)
print(f" train batch/device: {args.per_device_train_batch_size}", flush=True)
print(f" gradient_accumulation: {args.gradient_accumulation_steps}", flush=True)
print(f" effective batch size: {effective_batch_size}", flush=True)
@@ -1018,6 +1255,11 @@ def main() -> None:
vocab_size=vocab_size,
mlm_probability=args.mlm_probability,
special_token_ids=list(special_ids.values()),
+ masking_strategy=args.masking_strategy,
+ span_p_geom=args.span_p_geom,
+ span_max_length=args.span_max_length,
+ heteroatom_start_weight=args.heteroatom_start_weight,
+ ids_to_tokens=dict(tokenizer.ids_to_tokens),
)
report_to = [] if args.report_to == "none" else [args.report_to]
@@ -1080,6 +1322,14 @@ def main() -> None:
)
world_size = training_args.world_size if hasattr(training_args, "world_size") else 1
+
+ if args.masking_strategy in {"span", "hetero_span"} and args.num_workers < 2:
+ log(
+ "Warning: masking_strategy='span'/'hetero_span' runs in Python on the "
+ "data-loader path. Consider --num_workers >= 4 to overlap collation with "
+ "GPU compute and avoid becoming a training bottleneck."
+ )
+
log_training_plan(args, backend, n_params=n_params, world_size=world_size)
trainer = Trainer(
diff --git a/tests/conftest.py b/tests/conftest.py
index bdea6d9..5b07a2c 100644
--- a/tests/conftest.py
+++ b/tests/conftest.py
@@ -23,9 +23,7 @@ def find_existing_minimal_model() -> Path | None:
path.exists()
and (path / "config.json").exists()
and (path / "tokenizer.json").exists()
- and (
- any(path.glob("*.safetensors")) or (path / "pytorch_model.bin").exists()
- )
+ and (any(path.glob("*.safetensors")) or (path / "pytorch_model.bin").exists())
):
return path
return None
@@ -37,7 +35,6 @@ def existing_minimal_model() -> Path:
model = find_existing_minimal_model()
if model is None:
pytest.skip(
- "No existing minimal trained model found. "
- "Run a debug/smoke training command first."
+ "No existing minimal trained model found. Run a debug/smoke training command first."
)
return model
diff --git a/tests/test_collator.py b/tests/test_collator.py
index 690d6bd..bd322f5 100644
--- a/tests/test_collator.py
+++ b/tests/test_collator.py
@@ -91,9 +91,7 @@ def test_collator_random_replacements_never_use_special_ids():
replaced_random = (labels != -100) & (input_ids != labels) & (input_ids != 4)
if replaced_random.any():
- assert torch.all(
- ~torch.isin(input_ids[replaced_random], torch.tensor([0, 1, 2, 3, 4]))
- )
+ assert torch.all(~torch.isin(input_ids[replaced_random], torch.tensor([0, 1, 2, 3, 4])))
def test_collator_forces_at_least_one_mask_when_probability_nonzero():
diff --git a/tests/test_eval_moleculenet.py b/tests/test_eval_moleculenet.py
index 31f5fae..e6972de 100644
--- a/tests/test_eval_moleculenet.py
+++ b/tests/test_eval_moleculenet.py
@@ -297,9 +297,7 @@ def test_scaffold_split_raises_on_empty_valid_or_test() -> None:
}
)
- with pytest.raises(
- RuntimeError, match="Scaffold split produced an empty valid or test split"
- ):
+ with pytest.raises(RuntimeError, match="Scaffold split produced an empty valid or test split"):
split_sanitized_frame(
frame,
split="scaffold",
diff --git a/tests/test_eval_molformer.py b/tests/test_eval_molformer.py
index 810f374..9d75bcb 100644
--- a/tests/test_eval_molformer.py
+++ b/tests/test_eval_molformer.py
@@ -48,9 +48,7 @@ def test_molformer_registry_constructs_featurizer() -> None:
@pytest.mark.molformer
def test_molformer_embedding_smoke() -> None:
if not _molformer_enabled():
- pytest.skip(
- "Set MODERNMOLBERT_RUN_MOLFORMER_TESTS=1 to run MoLFormer smoke tests."
- )
+ pytest.skip("Set MODERNMOLBERT_RUN_MOLFORMER_TESTS=1 to run MoLFormer smoke tests.")
featurizer = HuggingFaceSmilesFeaturizer(
name="molformer_xl_both_10pct",
diff --git a/tests/test_span_masking.py b/tests/test_span_masking.py
new file mode 100644
index 0000000..73193b5
--- /dev/null
+++ b/tests/test_span_masking.py
@@ -0,0 +1,287 @@
+"""Tests for span and hetero_span masking strategies in MolecularMLMCollator."""
+
+import pytest
+import torch
+
+from modernmolbert.train_selfies_ape_modernbert import MolecularMLMCollator
+
+
+# ---------------------------------------------------------------------------
+# Shared fixtures
+# ---------------------------------------------------------------------------
+
+SPECIAL_IDS = [0, 1, 2, 3, 4] # bos, pad, eos, unk, mask
+VOCAB_SIZE = 100
+
+
+def _make_collator(
+ strategy: str = "standard",
+ mlm_probability: float = 0.15,
+ span_max_length: int = 6,
+ span_p_geom: float = 0.4,
+ heteroatom_start_weight: float = 2.0,
+ ids_to_tokens: dict | None = None,
+ vocab_size: int = VOCAB_SIZE,
+) -> MolecularMLMCollator:
+ if ids_to_tokens is None:
+ ids_to_tokens = {i: f"[T{i}]" for i in range(10, vocab_size)}
+ ids_to_tokens.update({0: "", 1: "", 2: "", 3: "", 4: ""})
+ return MolecularMLMCollator(
+ pad_token_id=1,
+ mask_token_id=4,
+ vocab_size=vocab_size,
+ mlm_probability=mlm_probability,
+ special_token_ids=SPECIAL_IDS,
+ masking_strategy=strategy,
+ span_p_geom=span_p_geom,
+ span_max_length=span_max_length,
+ heteroatom_start_weight=heteroatom_start_weight,
+ ids_to_tokens=ids_to_tokens,
+ )
+
+
+def _examples():
+ # bos=0, eos=2, real tokens 5-9
+ return [
+ {"input_ids": [0, 5, 6, 7, 8, 9, 2]},
+ {"input_ids": [0, 6, 7, 8, 2]},
+ {"input_ids": [0, 5, 6, 7, 8, 9, 5, 6, 7, 8, 2]},
+ ]
+
+
+# ---------------------------------------------------------------------------
+# __post_init__ validation
+# ---------------------------------------------------------------------------
+
+
+def test_invalid_span_p_geom_raises():
+ with pytest.raises(ValueError, match="span_p_geom"):
+ _make_collator("span", span_p_geom=0.0)
+
+
+def test_invalid_span_p_geom_one_raises():
+ with pytest.raises(ValueError, match="span_p_geom"):
+ _make_collator("span", span_p_geom=1.0)
+
+
+def test_invalid_span_max_length_raises():
+ with pytest.raises(ValueError, match="span_max_length"):
+ _make_collator("span", span_max_length=0)
+
+
+def test_invalid_heteroatom_start_weight_raises():
+ with pytest.raises(ValueError, match="heteroatom_start_weight"):
+ _make_collator("hetero_span", heteroatom_start_weight=0.0)
+
+
+def test_standard_strategy_no_validation():
+ # standard strategy ignores span params — should not raise
+ c = _make_collator("standard", span_p_geom=0.0, span_max_length=0)
+ assert c.masking_strategy == "standard"
+
+
+# ---------------------------------------------------------------------------
+# Output shape / dtype (same contract as standard)
+# ---------------------------------------------------------------------------
+
+
+@pytest.mark.parametrize("strategy", ["standard", "span", "hetero_span"])
+def test_output_shapes_match_standard(strategy):
+ torch.manual_seed(42)
+ std = _make_collator("standard", mlm_probability=0.3)
+ coll = _make_collator(strategy, mlm_probability=0.3)
+ examples = _examples()
+
+ std_batch = std(examples)
+ coll_batch = coll(examples)
+
+ for key in ("input_ids", "attention_mask", "labels"):
+ assert coll_batch[key].shape == std_batch[key].shape, f"{strategy}: shape mismatch on {key}"
+ assert coll_batch[key].dtype == std_batch[key].dtype, f"{strategy}: dtype mismatch on {key}"
+
+
+# ---------------------------------------------------------------------------
+# Invariants that must hold for all strategies
+# ---------------------------------------------------------------------------
+
+
+@pytest.mark.parametrize("strategy", ["standard", "span", "hetero_span"])
+def test_padding_positions_never_masked(strategy):
+ torch.manual_seed(0)
+ coll = _make_collator(strategy, mlm_probability=0.3)
+ batch = coll(_examples())
+ pad_positions = batch["attention_mask"] == 0
+ assert torch.all(batch["labels"][pad_positions] == -100), (
+ f"{strategy}: padding position appeared in labels"
+ )
+
+
+@pytest.mark.parametrize("strategy", ["standard", "span", "hetero_span"])
+def test_special_tokens_never_masked(strategy):
+ torch.manual_seed(0)
+ coll = _make_collator(strategy, mlm_probability=0.3)
+ batch = coll(_examples())
+ from torch.nn.utils.rnn import pad_sequence
+
+ examples = _examples()
+ ids = [torch.tensor(ex["input_ids"], dtype=torch.long) for ex in examples]
+ original = pad_sequence(ids, batch_first=True, padding_value=1)
+ special_mask = torch.zeros_like(original, dtype=torch.bool)
+ for sid in SPECIAL_IDS:
+ special_mask |= original.eq(sid)
+ assert torch.all(batch["labels"][special_mask] == -100), (
+ f"{strategy}: special token was included in masked positions"
+ )
+
+
+@pytest.mark.parametrize("strategy", ["standard", "span", "hetero_span"])
+def test_at_least_one_position_masked_per_batch(strategy):
+ torch.manual_seed(7)
+ coll = _make_collator(strategy, mlm_probability=0.3)
+ batch = coll(_examples())
+ assert (batch["labels"] != -100).any(), f"{strategy}: no position was masked"
+
+
+# ---------------------------------------------------------------------------
+# span-specific: contiguous spans
+# ---------------------------------------------------------------------------
+
+
+def test_span_produces_contiguous_runs():
+ """At least some masked runs in span mode should be length > 1."""
+ torch.manual_seed(0)
+ coll = _make_collator("span", mlm_probability=0.5, span_max_length=4)
+ # Use a long sequence to give spans room to form.
+ examples = [{"input_ids": [0] + list(range(5, 25)) + [2]}]
+ found_multi = False
+ for _ in range(20):
+ batch = coll(examples)
+ masked = batch["labels"][0] != -100
+ run_len = 0
+ for m in masked.tolist():
+ if m:
+ run_len += 1
+ if run_len >= 2:
+ found_multi = True
+ break
+ else:
+ run_len = 0
+ if found_multi:
+ break
+ assert found_multi, "span strategy never produced a run of length >= 2"
+
+
+def test_span_max_length_single_draw():
+ """With budget=1 and span_max_length=1, exactly 1 position is masked."""
+ for seed in range(10):
+ torch.manual_seed(seed)
+ # 1 BOS + 20 body + 1 EOS = 22 total; 20 eligible.
+ # budget = max(1, round(20 * 0.05)) = 1
+ coll = _make_collator("span", mlm_probability=0.05, span_max_length=1)
+ examples = [{"input_ids": [0] + list(range(5, 25)) + [2]}]
+ batch = coll(examples)
+ n_masked = int((batch["labels"] != -100).sum().item())
+ assert n_masked == 1, (
+ f"seed={seed}: expected 1 masked position with budget=1 and "
+ f"span_max_length=1, got {n_masked}"
+ )
+
+
+def test_span_max_length_clamps_individual_draw():
+ """Single-draw budget: max contiguous run ≤ span_max_length."""
+
+ def max_run(labels_row: torch.Tensor) -> int:
+ masked = (labels_row != -100).tolist()
+ best = cur = 0
+ for m in masked:
+ cur = cur + 1 if m else 0
+ best = max(best, cur)
+ return best
+
+ # 10 eligible positions, mlm_probability=0.05 → budget = max(1, round(10*0.05)) = 1
+ for max_len in [1, 2, 3, 4, 6]:
+ for seed in range(8):
+ torch.manual_seed(seed)
+ coll = _make_collator(
+ "span",
+ mlm_probability=0.05,
+ span_max_length=max_len,
+ span_p_geom=0.01, # very low p → geometric returns large pre-clamp values
+ )
+ examples = [{"input_ids": [0] + list(range(5, 15)) + [2]}]
+ batch = coll(examples)
+ run = max_run(batch["labels"][0])
+ assert run <= max_len, (
+ f"seed={seed}, span_max_length={max_len}: max run {run} exceeds span_max_length"
+ )
+
+
+# ---------------------------------------------------------------------------
+# hetero_span: weight tensor
+# ---------------------------------------------------------------------------
+
+
+def test_build_token_start_weights_special_tokens_zero():
+ ids_to_tokens = {i: f"[TOK{i}]" for i in range(VOCAB_SIZE)}
+ coll = _make_collator("hetero_span", ids_to_tokens=ids_to_tokens)
+ weights = coll._token_start_weights
+ assert weights is not None
+ for sid in SPECIAL_IDS:
+ assert weights[sid].item() == 0.0, f"special token {sid} has non-zero weight"
+
+
+def test_build_token_start_weights_heteroatom_elevated():
+ ids_to_tokens = {
+ 0: "",
+ 1: "",
+ 2: "",
+ 3: "",
+ 4: "",
+ 5: "[C]", # carbon — weight 1.0
+ 6: "[N]", # nitrogen — heteroatom, weight > 1.0
+ 7: "[O]", # oxygen — heteroatom, weight > 1.0
+ 8: "[Cl]", # chlorine — heteroatom, weight > 1.0
+ }
+ coll = _make_collator(
+ "hetero_span",
+ ids_to_tokens=ids_to_tokens,
+ vocab_size=9,
+ heteroatom_start_weight=2.0,
+ )
+ weights = coll._token_start_weights
+ assert weights is not None
+ assert weights[5].item() == 1.0, "[C] should have weight 1.0"
+ assert weights[6].item() == 2.0, "[N] should have heteroatom weight"
+ assert weights[7].item() == 2.0, "[O] should have heteroatom weight"
+ assert weights[8].item() == 2.0, "[Cl] should have heteroatom weight"
+
+
+def test_hetero_span_output_valid():
+ torch.manual_seed(1)
+ ids_to_tokens = {
+ **{sid: f"" for sid in SPECIAL_IDS},
+ **{i: "[N]" if i % 3 == 0 else "[C]" for i in range(5, VOCAB_SIZE)},
+ }
+ coll = _make_collator("hetero_span", ids_to_tokens=ids_to_tokens, mlm_probability=0.3)
+ batch = coll(_examples())
+ assert (batch["labels"] != -100).any()
+ assert torch.all(batch["labels"][batch["attention_mask"] == 0] == -100)
+
+
+# ---------------------------------------------------------------------------
+# Backwards compat: standard path unchanged
+# ---------------------------------------------------------------------------
+
+
+def test_standard_collator_unchanged():
+ """Existing standard collator with no new args still works."""
+ coll = MolecularMLMCollator(
+ pad_token_id=1,
+ mask_token_id=4,
+ vocab_size=32,
+ mlm_probability=0.3,
+ special_token_ids=SPECIAL_IDS,
+ )
+ torch.manual_seed(0)
+ batch = coll(_examples())
+ assert batch["input_ids"].shape == batch["labels"].shape