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
28 changes: 28 additions & 0 deletions electrolyte_fm/data_modules/isotope_dataset.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
from .property_prediction_dataset import PropertyPredictionDataModule
from .molnet_dataset import train_val_test_split
from datasets import Dataset, load_dataset
from pathlib import Path


class IsotopeDataModule(PropertyPredictionDataModule):
def __init__(self, path: str, **kwargs):
# Set default smi_column
self.path = path
assert Path(self.path).is_file()
super().__init__(**kwargs)

def prepare_data(self):
# Fetch data from the head node
self.dataset

def _get_dataset(self):
# Load the dataset
ds: Dataset = load_dataset(
"csv",
data_files=[self.path],
split="train",
keep_in_memory=False,
save_infos=False,
) # type: ignore

return train_val_test_split(ds)
86 changes: 86 additions & 0 deletions electrolyte_fm/utils/progressive_thawing.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,86 @@
import logging
from fnmatch import fnmatch
from lightning.pytorch.callbacks import BaseFinetuning


class ProgressiveThawing(BaseFinetuning):
def __init__(
self,
initial: list[str],
stages: list[list[str]],
stage_duration: int = 1,
lr_initial: float = 10,
lr_decay: float = 0.99,
):
super().__init__()
self.initial = initial
self.stages = stages
self.stage_duration = stage_duration
self.current_stage = 0
self.lr_scale = [0]
self.lr_initial = float(lr_initial)
self.lr_decay = float(lr_decay)

def setup(self, trainer, pl_module, stage) -> None:
if not hasattr(pl_module, "encoder"):
pl_module.configure_model()
self.freeze_before_training(pl_module)

@classmethod
def matching_modules(cls, pl_module, patterns):
for pattern in patterns:
pattern_matched = False
for name, module in pl_module.named_modules():
if name == "":
continue
logging.debug(
"Matching %s against %s: %d", name, pattern, fnmatch(name, pattern)
)
if fnmatch(name, pattern):
pattern_matched = True
yield name, module

if not pattern_matched:
names = [name for name, _ in pl_module.named_modules()]
raise ValueError(f"Pattern {pattern} not found in module: {names}")

def unfreeze_and_add_param_group(self, modules, optimizer):
self.make_trainable(modules)
params = self.filter_params(modules)
params = self.filter_on_optimizer(optimizer, params)
base_lr = optimizer.param_groups[0]["lr"]
if params:
self.lr_scale.append(self.lr_initial)
lr_scale = self.lr_scale[-1] + 1
optimizer.add_param_group({"params": params, "lr": base_lr / lr_scale})

def on_train_batch_start(self, trainer, *args, **kwargs):
# Decay lr offset
for idx in range(len(self.lr_scale)):
self.lr_scale[idx] *= self.lr_decay

for opt in trainer.optimizers:
base_lr = opt.param_groups[0]["lr"]
for idx, param_group in enumerate(opt.param_groups):
param_group["lr"] = base_lr / (self.lr_scale[idx] + 1)

def lr_scheduler_step(self, scheduler, metric) -> None:
print("LR Scheduler step", scheduler, metric)

def freeze_before_training(self, pl_module):
for name, module in self.matching_modules(pl_module, self.initial):
logging.info("Freezing %s", name)
self.freeze_module(module)

def finetune_function(self, pl_module, epoch: int, optimizer) -> None:
if self.current_stage == len(self.stages):
return

if epoch == self.stage_duration * (self.current_stage + 1):
thaw = self.stages[self.current_stage]
self.current_stage += 1
to_thaw = []
for name, module in self.matching_modules(pl_module, thaw):
logging.info("Thawing %s", name)
to_thaw.append(module)
self.unfreeze_and_add_param_group(to_thaw, optimizer)
112 changes: 112 additions & 0 deletions opt/interp_embeddings/embedding_figure.jl
Original file line number Diff line number Diff line change
@@ -1,5 +1,9 @@
using Makie
using DataFrames
using Random
using Statistics
using UMAP
using Distances
using CairoMakie: CairoMakie
using CSV: CSV
using CategoricalArrays: categorical, levelcode
Expand Down Expand Up @@ -169,3 +173,111 @@ function figure_olfactory()

return f
end


function classify_decay_from_NZ(N::Real, Z::Real;
tol_offset::Real=0.8, tol_scale::Real=0.02)

# Beta-stability valley (from SEMF): Z_beta(A) ≈ A / (2 + 0.015 * A^(2/3))
predict_Z_beta(A::Real) = A / (2 + 0.015 * A^(2/3))

A = N + Z
# very heavy nuclei
if Z ≥ 92 && A ≥ 240
return "Spontaneous \nFission"
elseif Z ≥ 84 && A ≥ 210
return L"$\alpha$"
end
zβ = predict_Z_beta(A)
δ = Z - zβ
tol = tol_offset + tol_scale * Z # widen tolerance for heavier Z
if abs(δ) ≤ tol
return "Stable"
elseif δ < 0
return L"$\beta^-$" # too many neutrons -> beta- decay
else
return L"$\beta^+$" # too many protons -> beta+/EC
end
end

function figure_isotopes_umap(n_neighbors::Int=10, min_dist::Real=2, metric::Symbol=:manhattan, seed::Int=42)

csv_path = "isotope_embeddings.csv"
df = DataFrame(CSV.File(csv_path))

# Labels from N/Z
N = Float64.(coalesce.(df.neutrons, NaN))
Z = Float64.(coalesce.(df.protons, NaN))
labels_raw = [classify_decay_from_NZ(N[i], Z[i]) for i in eachindex(N)]

# Embedding matrix: numeric columns excluding smiles/N/Z
exclude = Set([:smiles, :neutrons, :protons])
_is_numeric_col(col) = (eltype(col) <: Real) || (eltype(col) <: Union{Missing,Real})
embed_cols = [nm for nm in names(df) if nm ∉ exclude && _is_numeric_col(df[!, nm])]
isempty(embed_cols) && error("No numeric embedding columns found (after excluding smiles/N/Z).")

X = Matrix{Float64}(coalesce.(df[!, embed_cols], 0.0))

# Standardize columns
for j in axes(X, 2)
μ, σ = mean(X[:, j]), std(X[:, j])
if σ > 0 && isfinite(σ)
X[:, j] .= (X[:, j] .- μ) ./ σ
else
X[:, j] .= 0.0
end
end

# UMAP → 2D
Random.seed!(seed)
metric_obj = if metric == :cosine
CosineDist()
elseif metric == :manhattan
Cityblock()
else
Euclidean()
end
Y = umap(Matrix(X'), 2; n_neighbors=n_neighbors, min_dist=min_dist, metric=metric_obj)
Y = Matrix(Y')

canonical = [L"$\beta^+$", L"$\beta^-$", L"$\alpha$", "Stable", "Spontaneous Fission"]
present = unique(labels_raw)
levels = [c for c in canonical if c in present]
append!(levels, [x for x in present if x ∉ Set(canonical)])
decay_cat = categorical(labels_raw; levels=levels, ordered=true)
n = length(levels)

f = Figure(;size = (1.5inch, 1.5inch))
ax = Axis(
f[1, 1];
limits = (nothing, (minimum(Y[:, 2]) - 7, nothing)),
)
hidedecorations!(ax)

palette = MISTStyle.CAT_COLORS[1:n]

h = scatter!(ax, Y[:, 1], Y[:, 2];
color = levelcode.(decay_cat),
colormap = palette,
colorrange = (1, n),
marker = :circle,
)

# Create legend elements
elements = [MarkerElement(;
marker = :circle,
color = palette[i],
) for i in 1:length(levels)]

Legend(f[1, 1], elements, levels;
tellheight = false,
tellwidth = false,
orientation = :horizontal,
halign = :right,
valign = :bottom,
padding=(1, 1, 1, 1),
margin=(1, 1, 1, 1),
)

return f
end
77 changes: 77 additions & 0 deletions submit/isotopes.jsonnet
Original file line number Diff line number Diff line change
@@ -0,0 +1,77 @@
{
walltime: '1:0:0',
nodes: 1, // Multi-node is not currently supported. Config is only on leader node
env: {
TOKENIZERS_PARALLELISM: true,
},
train: {
tags: ['finetuning', 'isotopes', 'thaw'],
data: {
class_path: 'electrolyte_fm.data_modules.IsotopeDataModule',
init_args: {
path: '/nfs/turbo/coe-venkvis/abhutani/electrolyte-fm/opt/isotope_half_lives.csv',
batch_size: 16,
val_batch_size: 2 * self.batch_size,
tokenizer: $.train.model.init_args.encoder_ckpt,
target_columns: ['half_life', 'log_time'],
num_workers: 4,
prefetch_factor: 8,
},
},
model: {
class_path: 'electrolyte_fm.models.LMFinetuning',
init_args: {
encoder_ckpt: '/nfs/turbo/coe-venkvis/mist/ti624ev1/pretrained/checkpoints/last.ckpt',
task: 'regression',
metrics: ['mae', 'mae-channel', 'r2-channel'],
freeze_encoder: false,
transform: ['standardize', 'standardize'],
output_size: std.length($.train.data.init_args.target_columns),
target_columns: $.train.data.init_args.target_columns,

// Duplicate pre-training optimizer config
optimizer: {
class_path: 'torch.optim.AdamW',
init_args: {
lr: 1.6e-4,
weight_decay: 0.01,
},
},

lr_schedule: {
class_path: 'electrolyte_fm.utils.lr_schedule.RelativeCosineWarmup',
init_args: {
num_training_steps: $.train.trainer.max_steps,
num_warmup_steps: 'beta2',
rel_decay: 0.1,
},
},
},
},
trainer: {
max_steps: 5000,
precision: 'bf16-true',
enable_progress_bar: false,
strategy: 'auto',
callbacks: [
{
class_path: 'electrolyte_fm.utils.progressive_thawing.ProgressiveThawing',
init_args: {
initial: ['encoder'],
stages: [
['encoder.embeddings'],
['encoder.encoder.layer.7'],
['encoder.encoder.layer.6'],
['encoder.encoder.layer.5'],
['encoder.encoder.layer.4'],
['encoder.encoder.layer.3'],
['encoder.encoder.layer.2'],
['encoder.encoder.layer.1'],
],
stage_duration: 3,
},
},
],
},
},
}
Loading