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
30 changes: 25 additions & 5 deletions analysis/sweep/fixed_eval_best_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -160,6 +160,17 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--max_seq_length", type=int, default=128)
parser.add_argument("--expected_vocab_size", type=int, default=631)
parser.add_argument("--expected_max_position_embeddings", type=int, default=128)
parser.add_argument(
"--representation",
choices=["SELFIES", "SMILES"],
default="SELFIES",
help="Representation of the swept checkpoints (selects the default validation column).",
)
parser.add_argument(
"--molecule_column",
default=None,
help="Validation parquet column (default: selfies / smiles_canonical_clean by representation).",
)
parser.add_argument(
"--limit",
type=int,
Expand Down Expand Up @@ -323,11 +334,13 @@ def load_tokenizer(tokenizer_dir: Path):
return AutoTokenizer.from_pretrained(str(tokenizer_dir), trust_remote_code=True)


def load_valid_full(valid_parquet: Path, *, limit: int | None) -> list[str]:
def load_valid_full(
valid_parquet: Path, *, limit: int | None, molecule_column: str = "selfies"
) -> list[str]:
if not valid_parquet.exists():
raise FileNotFoundError(f"Missing validation parquet: {valid_parquet}")
frame = pd.read_parquet(valid_parquet, columns=["selfies"])
seqs = [str(value).strip() for value in frame["selfies"] if str(value).strip()]
frame = pd.read_parquet(valid_parquet, columns=[molecule_column])
seqs = [str(value).strip() for value in frame[molecule_column] if str(value).strip()]
return seqs[:limit] if limit is not None else seqs


Expand All @@ -337,6 +350,7 @@ def load_valid_train_matched(
n_examples: int,
seed: int,
shuffle_buffer_size: int,
molecule_column: str = "selfies",
) -> list[str]:
ds = load_dataset(
"parquet",
Expand All @@ -347,7 +361,7 @@ def load_valid_train_matched(
ds = ds.shuffle(seed=seed, buffer_size=shuffle_buffer_size)
seqs: list[str] = []
for row in ds:
seq = str(row.get("selfies", "")).strip()
seq = str(row.get(molecule_column, "")).strip()
if not seq:
continue
seqs.append(seq)
Expand Down Expand Up @@ -605,14 +619,20 @@ def main() -> None:
ids_to_tokens = dict(getattr(tokenizer, "ids_to_tokens", {}))
log.info(" vocab_size=%d pad=%d mask=%d", vocab_size, pad_token_id, mask_token_id)

molecule_column = args.molecule_column or (
"smiles_canonical_clean" if args.representation == "SMILES" else "selfies"
)
train_matched_n = min(4096, args.limit) if args.limit is not None else 4096
eval_sets = {
"valid_full": load_valid_full(args.valid_parquet, limit=args.limit),
"valid_full": load_valid_full(
args.valid_parquet, limit=args.limit, molecule_column=molecule_column
),
"valid_4096_train_matched": load_valid_train_matched(
args.valid_parquet,
n_examples=train_matched_n,
seed=242,
shuffle_buffer_size=100_000,
molecule_column=molecule_column,
),
}
for name, seqs in eval_sets.items():
Expand Down
11 changes: 11 additions & 0 deletions configs/featurizers/modernmolbert_smiles.json
Original file line number Diff line number Diff line change
@@ -0,0 +1,11 @@
{
"batch_size": 32,
"device": "auto",
"max_seq_length": 128,
"model_dir": "runs/best_chembl36_smiles_small/final_model",
"name": "modernmolbert_smiles",
"pooling": "mean",
"representation": "SMILES",
"tokenizer_path": "runs/best_chembl36_smiles_small/final_model",
"type": "modernmolbert_smiles"
}
80 changes: 70 additions & 10 deletions scripts/sweeps/run_sweep.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,12 +22,31 @@
from pathlib import Path

# ─── Fixed across all runs ────────────────────────────────────────────────────
# The prepared chembl36 dataset carries both a `selfies` and a `smiles_canonical_clean`
# column, so both representations reuse the same dataset directory read-only.
DATASET_DIR = "data/pretrain/chembl36_selfies"
SELFIES_COLUMN = "selfies"
TRAIN_SPLIT = "train"
VALIDATION_SPLIT = "valid"
TOKENIZER_PATH = "tokenizer/chembl36_selfies_2m_ape_max2_min3000.json"
TOKENIZER_METADATA_PATH = "tokenizer/chembl36_selfies_2m_ape_max2_min3000.metadata.json"

# Per-representation tokenizer, dataset column, and masking grid. SELFIES keeps the
# historical defaults; SMILES points at the SMILES APE tokenizer and drops hetero_span
# (its heteroatom bias is SELFIES-bracket-specific and degrades to plain span on SMILES).
REPRESENTATION_DEFAULTS = {
"SELFIES": {
"tokenizer_path": "tokenizer/chembl36_selfies_2m_ape_max2_min3000.json",
"tokenizer_metadata_path": "tokenizer/chembl36_selfies_2m_ape_max2_min3000.metadata.json",
"molecule_column": "selfies",
"masking": ["standard", "span", "hetero_span"],
"run_root_tag": "",
},
"SMILES": {
"tokenizer_path": "tokenizer/chembl36_smiles_2m_ape_max6_mf3000.json",
"tokenizer_metadata_path": "tokenizer/chembl36_smiles_2m_ape_max6_mf3000.metadata.json",
"molecule_column": "smiles_canonical_clean",
"masking": ["standard", "span"],
"run_root_tag": "smiles_",
},
}

MAX_SEQ_LENGTH = 128
MAX_STEPS = 30000
Expand Down Expand Up @@ -69,12 +88,21 @@ def parse_args() -> argparse.Namespace:
description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter
)
parser.add_argument("--model-size", choices=sorted(PRESETS), required=True)
parser.add_argument(
"--representation",
choices=sorted(REPRESENTATION_DEFAULTS),
default="SELFIES",
help="Molecular string representation (default: SELFIES).",
)
parser.add_argument(
"--masking",
nargs="+",
choices=ALL_MASKING,
default=ALL_MASKING,
help="Masking strategies to sweep (default: all three).",
default=None,
help=(
"Masking strategies to sweep (default: per representation; "
"SELFIES uses all three, SMILES uses standard+span)."
),
)
parser.add_argument(
"--mlm-probs",
Expand All @@ -96,8 +124,21 @@ def parse_args() -> argparse.Namespace:
help="Output root (default: runs/chembl36_<model-size>_mask_mlm_lr_sweep).",
)
parser.add_argument("--dataset-dir", default=DATASET_DIR)
parser.add_argument("--tokenizer-path", default=TOKENIZER_PATH)
parser.add_argument("--tokenizer-metadata-path", default=TOKENIZER_METADATA_PATH)
parser.add_argument(
"--molecule-column",
default=None,
help="Dataset column (default: per representation).",
)
parser.add_argument(
"--tokenizer-path",
default=None,
help="Tokenizer vocabulary JSON (default: per representation).",
)
parser.add_argument(
"--tokenizer-metadata-path",
default=None,
help="Tokenizer metadata JSON (default: per representation).",
)
parser.add_argument(
"--dry-run",
action="store_true",
Expand Down Expand Up @@ -144,8 +185,10 @@ def build_command(
"modernmolbert.train_selfies_ape_modernbert",
"--dataset_name",
args.dataset_dir,
"--selfies_column",
SELFIES_COLUMN,
"--representation",
args.representation,
"--molecule_column",
args.molecule_column,
"--train_split",
TRAIN_SPLIT,
"--use_validation_split",
Expand Down Expand Up @@ -222,7 +265,24 @@ def main() -> None:
args = parse_args()
preset = PRESETS[args.model_size]
learning_rates = args.learning_rates or preset["learning_rates"]
run_root = args.run_root or Path(f"runs/chembl36_{args.model_size}_mask_mlm_lr_sweep")

# Fill representation-dependent defaults for anything not overridden on the CLI.
rep = REPRESENTATION_DEFAULTS[args.representation]
args.tokenizer_path = args.tokenizer_path or rep["tokenizer_path"]
args.tokenizer_metadata_path = args.tokenizer_metadata_path or rep["tokenizer_metadata_path"]
args.molecule_column = args.molecule_column or rep["molecule_column"]
args.masking = args.masking or rep["masking"]

invalid = [m for m in args.masking if m not in rep["masking"]]
if invalid:
sys.exit(
f"ERROR: masking {invalid} not valid for representation {args.representation}; "
f"allowed: {rep['masking']}"
)

run_root = args.run_root or Path(
f"runs/chembl36_{rep['run_root_tag']}{args.model_size}_mask_mlm_lr_sweep"
)

if not args.dry_run:
preflight(args)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,12 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--device", default="auto")
parser.add_argument("--max-seq-length", type=int, default=256)
parser.add_argument("--pooling", choices=["mean", "cls"], default="mean")
parser.add_argument(
"--representation",
choices=["SELFIES", "SMILES"],
default="SELFIES",
help="Checkpoint's molecular representation (SMILES skips SMILES->SELFIES conversion).",
)
parser.add_argument("--overwrite", action=argparse.BooleanOptionalAction, default=False)
return parser.parse_args()

Expand Down Expand Up @@ -101,6 +107,7 @@ def make_featurizer(args: argparse.Namespace):
pooling=args.pooling,
device=args.device,
batch_size=args.batch_size,
representation=args.representation,
)


Expand Down
Loading
Loading