Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
56 commits
Select commit Hold shift + click to select a range
3f17cf8
migrate TokenizerStats to uv
awadell1 May 2, 2025
f0f1ece
fix: ssl error causing download issues
awadell1 May 6, 2025
cded7fa
fixes for running tokenizer stats with uv
awadell1 May 6, 2025
8f7a8ba
tweak: compute loss using finetuned n-gram model as well
awadell1 May 6, 2025
4046bad
fix: correct encoding for ChemGPT
awadell1 May 6, 2025
f036707
Summarize results with fixed effects
awadell1 May 6, 2025
d6fbae0
move plot code to a script
awadell1 May 6, 2025
9481dde
add script for finding ambiguous smiles
awadell1 May 7, 2025
93a5527
expand fe models to cover finetuned models
awadell1 May 7, 2025
7876c60
add script for compressing tok stats
awadell1 May 7, 2025
dc30fcf
feat: add fixed effects models comparing fits
awadell1 May 8, 2025
8f2a1e1
add fixed effects figure
awadell1 May 8, 2025
7248875
refactor fe figure, add more scale options
awadell1 May 8, 2025
d9177ea
add precompile flag to submit jobs
awadell1 May 8, 2025
6d5bd93
update figure
awadell1 May 8, 2025
9340431
update deps and add intrinsic plots
awadell1 May 8, 2025
1132d60
Add 90% percentile span to intrinsic metric plots
awadell1 May 9, 2025
519bb60
update plots
awadell1 May 9, 2025
b83b137
add ngram statistics plot
awadell1 May 9, 2025
c8e97ec
fixup and save tokenizer stats
awadell1 May 9, 2025
839fb3e
fix barploterrors
awadell1 May 9, 2025
687dae5
fix typo
awadell1 May 9, 2025
bc4f55d
plot tweaks
awadell1 May 10, 2025
cda20bc
refactor: dedicated tokenizer dataset loading
awadell1 May 10, 2025
eedcad8
fixup typos
awadell1 May 10, 2025
039d8e0
fix spe and expand tokenizer summary table
awadell1 May 12, 2025
f08c03e
don't archive unmerged files
awadell1 May 12, 2025
2e88bef
fix: replace paths with real ones when loading
awadell1 May 12, 2025
a94f0a2
fix paths in archive
awadell1 May 12, 2025
e79a73b
fix typo in archive
awadell1 May 12, 2025
ce866f8
misc fixes
awadell1 May 18, 2025
00de458
more tests and benchmark loader
awadell1 May 21, 2025
5f4805b
manage batch jobs in submit_jobs
awadell1 May 21, 2025
d266329
switch to UInt32 for token ids and squash bug
awadell1 May 22, 2025
ec2f8e1
fix oov launcher
awadell1 May 22, 2025
cff438b
feat: resume partially completed job arrays
awadell1 May 22, 2025
4515254
submit all jobs
awadell1 May 22, 2025
c83740d
support UInt32 for MaskedCode
awadell1 May 22, 2025
fbc9cda
reduce log rate
awadell1 May 22, 2025
df48d84
squash some bugs
awadell1 May 22, 2025
66c532e
feat: switch dgx to apptainer
awadell1 Feb 23, 2025
a45e300
feat: batched async molecular transcoding
awadell1 Jun 2, 2025
87eae38
feat: add randomization support to pretraining
awadell1 Jun 2, 2025
396ddd2
feat: tracy support
awadell1 Jun 3, 2025
5b397a1
fix: avoid converting ngrams to dict during serialization
awadell1 Jun 2, 2025
2ad4d75
feat: fix merge and refactor to reduce memory usage
awadell1 Jun 2, 2025
693e99f
add tracy profile option
awadell1 Jun 13, 2025
a95ed3d
feat: transfer tokenizer stats archive to dataden
awadell1 Jun 13, 2025
18640c0
fix pre-commit and import error
awadell1 Sep 1, 2025
ebfa2e7
merge updates from main
awadell1 Sep 1, 2025
87e389b
Merge branch 'master' into smirk-paper-plots
anoushka2000 Oct 21, 2025
634160e
fix: pre-commit
anoushka2000 Oct 21, 2025
3fcd19d
mark shebang executable
anoushka2000 Oct 21, 2025
a9db85a
formatting
anoushka2000 Oct 21, 2025
eae897f
Merge remote-tracking branch 'origin/master' into smirk-paper-plots
awadell1 Oct 21, 2025
f13512c
formatting
anoushka2000 Oct 21, 2025
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
71 changes: 2 additions & 69 deletions electrolyte_fm/data_modules/molnet_dataset.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,7 @@
import logging

from datasets import Dataset, DatasetDict, load_dataset
from rdkit.Chem.Scaffolds.MurckoScaffold import MurckoScaffoldSmiles
from sklearn.model_selection import GroupShuffleSplit
from datasets import Dataset, load_dataset

from .property_prediction_dataset import PropertyPredictionDataModule
from .utils import scaffold_split, train_val_test_split

_URLS = {
"qm8": "https://deepchemdata.s3-us-west-1.amazonaws.com/datasets/qm8.csv",
Expand Down Expand Up @@ -87,67 +84,3 @@ def _get_dataset(self):
return ds
else:
raise ValueError(f"Unknown split {self.split}")


def train_val_test_split(ds, **kwargs):
ds_train_other = ds.train_test_split(test_size=0.2, seed=42, **kwargs)
ds_val_test = ds_train_other["test"].train_test_split(
test_size=0.5, seed=42, **kwargs
)
return DatasetDict(
{
"train": ds_train_other["train"],
"validation": ds_val_test["train"],
"test": ds_val_test["test"],
}
)


def scaffold_hash(smi: str) -> str:
try:
scaffold = MurckoScaffoldSmiles(smi)
except ValueError:
logging.warn("No scaffold for %s, using input smiles string", smi)
scaffold = smi
return scaffold


def scaffold_split(ds: Dataset, smi_column):
# Hash scaffolds and then bin into groups, maintains the scaffold split
# but reduces the compute
df = ds.map(
lambda x: {"scaffold": scaffold_hash(x)},
input_columns=smi_column,
batched=False,
).to_pandas(batched=False)

# Split
train, other = next(
GroupShuffleSplit(n_splits=1, test_size=0.2, random_state=42).split(
df.index, groups=df["scaffold"].values
)
)
df_other = df.iloc[other]
val, test = next(
GroupShuffleSplit(n_splits=1, test_size=0.5, random_state=42).split(
df_other, groups=df.iloc[other]["scaffold"]
)
)
return DatasetDict(
{
"train": Dataset.from_pandas(df.iloc[train], preserve_index=False),
"validation": Dataset.from_pandas(df_other.iloc[val], preserve_index=False),
"test": Dataset.from_pandas(df_other.iloc[test], preserve_index=False),
}
)


def strip_unk_tokens(encoding: dict, unk_token_id: int) -> dict:
"""Remove unknown tokens from input"""
is_oov = [id == unk_token_id for id in encoding["input_ids"]]
out = {}
for k, v in encoding.items():
assert len(v) == len(is_oov)
out[k] = [x for x, oov in zip(v, is_oov) if not oov]
out["is_oov"] = any(is_oov)
return out
4 changes: 3 additions & 1 deletion electrolyte_fm/data_modules/roberta_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@ def __init__(
persistent_workers=False,
canonical: Optional[bool] = None, # Deprecated: Use encoding instead
encoding: str | MolEncoding = "smiles",
random: bool = False,
):
super().__init__()

Expand All @@ -48,6 +49,7 @@ def __init__(
self.prefetch_factor = prefetch_factor
self.persistent_workers = persistent_workers
self.encoding = MolEncoding(encoding)
self.random = random
self.hparams["vocab_size"] = self.vocab_size
self.save_hyperparameters(logger=False)

Expand Down Expand Up @@ -85,7 +87,7 @@ def setup(self, stage: str) -> None:
)

# Transcode
ds = encode_molecules(ds, "text", encoding=self.encoding)
ds = encode_molecules(ds, "text", encoding=self.encoding, random=self.random)

# Tokenize
ds = ds.map(
Expand Down
87 changes: 83 additions & 4 deletions electrolyte_fm/data_modules/utils.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,15 @@
import logging
import asyncio
from enum import Enum
import random
from asyncio import Semaphore
from typing import TypeVar
import torch
from rdkit import Chem
from datasets import Dataset, DatasetDict, IterableDatasetDict
from datasets.distributed import split_dataset_by_node
from rdkit.Chem.Scaffolds.MurckoScaffold import MurckoScaffoldSmiles
from sklearn.model_selection import GroupShuffleSplit


def is_fast(tokenizer):
Expand Down Expand Up @@ -102,6 +106,7 @@ def encode_molecules(
output_column: str | None = None,
encoding: MolEncoding = MolEncoding.SMILES,
random: bool = False,
max_workers: int = 8,
**kwargs,
) -> AbstractDataset:
"""Convert SMILES encoding in `input_column` to desired `encoding` and save to `output_column`.
Expand All @@ -110,15 +115,89 @@ def encode_molecules(
"""
assert isinstance(input_column, str)
output_column = output_column or input_column

encode = encoding if not random else encoding.random

tasks = Semaphore(max_workers)

async def async_encode(batch: list[str]) -> dict:
async with tasks:
return {output_column: [encode(smi) for smi in batch]}

async def async_filter(batch: list[str | None]) -> list[bool]:
async with tasks:
return [x is not None for x in batch]

ds = ds.map(
lambda smi: {output_column: encode(smi)},
async_encode,
input_columns=input_column,
batched=False,
batched=True,
**kwargs,
)
return ds.filter(lambda x: x[output_column] is not None, batched=False, **kwargs)
return ds.filter(async_filter, batched=True, input_columns=output_column, **kwargs)


def train_val_test_split(ds, **kwargs):
ds_train_other = ds.train_test_split(test_size=0.2, seed=42, **kwargs)
ds_val_test = ds_train_other["test"].train_test_split(
test_size=0.5, seed=42, **kwargs
)
return DatasetDict(
{
"train": ds_train_other["train"],
"validation": ds_val_test["train"],
"test": ds_val_test["test"],
}
)


def scaffold_hash(smi: str) -> str:
try:
scaffold = MurckoScaffoldSmiles(smi)
except ValueError:
logging.warning("No scaffold for %s, using input smiles string", smi)
scaffold = smi
return scaffold


def scaffold_split(ds: Dataset, smi_column):
# Hash scaffolds and then bin into groups, maintains the scaffold split
# but reduces the compute
df = ds.map(
lambda x: {"scaffold": scaffold_hash(x)},
input_columns=smi_column,
batched=False,
).to_pandas(batched=False)

# Split
train, other = next(
GroupShuffleSplit(n_splits=1, test_size=0.2, random_state=42).split(
df.index, groups=df["scaffold"].values
)
)
df_other = df.iloc[other]
val, test = next(
GroupShuffleSplit(n_splits=1, test_size=0.5, random_state=42).split(
df_other, groups=df.iloc[other]["scaffold"]
)
)
return DatasetDict(
{
"train": Dataset.from_pandas(df.iloc[train], preserve_index=False),
"validation": Dataset.from_pandas(df_other.iloc[val], preserve_index=False),
"test": Dataset.from_pandas(df_other.iloc[test], preserve_index=False),
}
)


def strip_unk_tokens(encoding: dict, unk_token_id: int) -> dict:
"""Remove unknown tokens from input"""
is_oov = [id == unk_token_id for id in encoding["input_ids"]]
out = {}
for k, v in encoding.items():
assert len(v) == len(is_oov)
out[k] = [x for x, oov in zip(v, is_oov) if not oov]
out["is_oov"] = any(is_oov)
return out


def stack_columns(batch, columns: list[str], output: str, dtype=None):
Expand Down
13 changes: 10 additions & 3 deletions electrolyte_fm/utils/cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,17 +42,24 @@ def cached_github_archive(repo, commit, file):
return cached_path


def cached_download(url: str, path: Path) -> Path:
def cached_download(url: str, path: Path, disable_ssl=False) -> Path:
path = Path(path)
cache = Path(__file__).parent.parent.parent.joinpath(".cache")
cached_file = cache.joinpath(path)
cached_file.parent.mkdir(exist_ok=True, parents=True)
if not cached_file.exists():
import urllib
import ssl
import urllib.request

ctx = (
ssl.create_default_context()
if not disable_ssl
else ssl._create_unverified_context()
)

user_agent = "Wget/1.19.5" # Pretend to be wget
req = urllib.request.Request(url, headers={"User-Agent": user_agent})
with urllib.request.urlopen(req) as fid:
with urllib.request.urlopen(req, context=ctx) as fid:
cached_file.parent.mkdir(parents=True, exist_ok=True)
with open(cached_file, "wb") as out:
out.write(fid.read())
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from transformers import PreTrainedTokenizerBase
from transformers.tokenization_utils_base import BatchEncoding

from ..utils.cache import cached_download
from .cache import cached_download


class PreTrainedSPETokenizer(PreTrainedTokenizerBase):
Expand Down
7 changes: 5 additions & 2 deletions electrolyte_fm/utils/tokenizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
def load_tokenizer(name: str, **kwargs) -> PreTrainedTokenizerBase:
# Locate Tokeniser and dataset
unk_name = RuntimeError(f"Unknown tokenizer: {name}")
if name.startswith("smirk"):
if name in ["smirk", "smirk-selfies", "smirk-cls"]:
from smirk import SmirkTokenizerFast

if name == "smirk":
Expand All @@ -24,7 +24,7 @@ def load_tokenizer(name: str, **kwargs) -> PreTrainedTokenizerBase:
raise unk_name

elif name == "SmilesPE/SPE_ChEMBL":
from ..tokenize.spe import pretrained_spe_tokenizer
from .spe import pretrained_spe_tokenizer

return pretrained_spe_tokenizer(cache_generated=True)

Expand Down Expand Up @@ -319,6 +319,9 @@ def all_special_ids(self) -> list[int]:
"SmirkTokenizer", fast_tokenizer_class=SmirkTokenizerFast
)

if Path(name).exists():
name = str(Path(name).resolve())

tok_tf = AutoTokenizer.from_pretrained(
name,
trust_remote_code=True,
Expand Down
4 changes: 4 additions & 0 deletions opt/TokenizerStats/.gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -9,3 +9,7 @@ models
smirk-gpe-*
smirk-gpe-*/
fig/
*.tar.*
archive/
*.slurm
*.jld2
1 change: 1 addition & 0 deletions opt/TokenizerStats/.python-version
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
3.12
4 changes: 2 additions & 2 deletions opt/TokenizerStats/Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -12,13 +12,13 @@ LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e"
LogExpFunctions = "2ab3a3ac-af41-5b50-aa03-7779005ae688"
MPI = "da04e1cc-30fd-572f-bb4f-1f8673147195"
MPIPreferences = "3da0fdf6-3ccc-4f1b-acd9-58baa6c99267"
NVTX = "5da4648a-3479-48b8-97b9-01cb529c0a1f"
OnlineStats = "a15396b6-48d5-5d58-9928-6d29437db91e"
PythonCall = "6099a3de-0909-46bc-b1f4-468b9a2dfc0d"
SHA = "ea8e919c-243c-51af-8825-aaa63cd721ce"
Serialization = "9e88b42a-f829-5b0c-bbe9-9e923198166b"
SparseArrays = "2f01184e-e22b-5df5-ae63-d93ebab69eaf"
StatsBase = "2913bbd2-ae8a-5f71-8c99-4fb6c76f3a91"
Tracy = "e689c965-62c8-4b79-b2c5-8359227902fd"

[compat]
ArgParse = "1.2.0"
Expand All @@ -27,12 +27,12 @@ JSON = "0.21"
LinearAlgebra = "1.10"
MPI = "0.20"
MPIPreferences = "0.1"
NVTX = "0.3.4"
OnlineStats = "1.7"
PythonCall = "0.9"
Serialization = "1.10"
SparseArrays = "1.10"
StatsBase = "0.34"
Tracy = "0.1.4"

[extras]
ReTestItems = "817f1d60-ba6b-4fd5-9520-3cf149f6a823"
Expand Down
6 changes: 4 additions & 2 deletions opt/TokenizerStats/activate
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,10 @@ DIR="$(git rev-parse --show-toplevel)/opt/TokenizerStats"
# Load modules
if command -v module > /dev/null; then
module purge
module --ignore_cache load gcc python/3.11.5 openmpi/4.1.6
module --ignore_cache load gcc/10.3.0 openmpi/4.1.6 python/3.11.5
export SSL_CERT_DIR=/etc/pki/tls/certs
export SSL_CERT_FILE=/etc/pki/tls/cert.pem
export JULIA_CPU_TARGET="generic;znver4,clone_all;znver3,clone_all;haswell"
fi

# Activate virtual environment
Expand All @@ -16,6 +19,5 @@ export HF_HOME="$(git rev-parse --show-toplevel)/.cache/huggingface"
export TOKENIZERS_PARALLELISM=false

# Configure julia
export JULIA_CPU_TARGET="generic;znver4,clone_all;znver3,clone_all;haswell"
export JULIA_CONDAPKG_BACKEND=Null
export JULIA_PYTHONCALL_EXE="$DIR/.venv/bin/python"
13 changes: 10 additions & 3 deletions opt/TokenizerStats/benchmark/bench_tokenizer.jl
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
using BenchmarkTools
using TokenizerStats
using TokenizerStats: DatasetConfig, dataset_split
using PythonCall: pyconvert

const suite = BenchmarkGroup()
Expand Down Expand Up @@ -28,7 +28,14 @@ TOKENIZERS = [
"meta-llama/Meta-Llama-3-8B",
]

setup_dataset(tokenizer) = Iterators.take(TokenizerStats.molnet("freesolv"; tokenizer).train_dataset, 10)
function setup_dataset(tokenizer, encoding)
dc = DatasetConfig("qm9", tokenizer, encoding)
ds = dataset_split(dc, "val")
return Iterators.take(ds, 1000)
end
for name in TOKENIZERS
suite[name] = @benchmarkable map(obs -> pyconvert(Vector{Int}, obs["input_ids"]), ds) setup=(ds = setup_dataset($name))
suite[name] = tok_suite = BenchmarkGroup()
for encoding in ["smiles", "smiles-canonical", "smiles-kekule"]
tok_suite[encoding] = @benchmarkable foreach(identity, ds) setup=(ds = setup_dataset($name, $encoding))
end
end
Loading