Feature Extraction
Transformers
Safetensors
English
mist_ionic_conductivity
mist
chemistry
molecular-property-prediction
custom_code
Instructions to use mist-models/mist-conductivity-27.0M-2mpg8dcd with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use mist-models/mist-conductivity-27.0M-2mpg8dcd with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="mist-models/mist-conductivity-27.0M-2mpg8dcd", trust_remote_code=True)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("mist-models/mist-conductivity-27.0M-2mpg8dcd", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
| import importlib | |
| import json | |
| import logging | |
| import math | |
| import os | |
| from math import cos, pi | |
| from pathlib import Path | |
| from typing import Any, Callable, Dict, List, Optional, Tuple, Union | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from datasets import IterableDataset | |
| from deepspeed.utils.zero_to_fp32 import \ | |
| get_fp32_state_dict_from_zero_checkpoint | |
| from jsonargparse import Namespace | |
| from lightning.pytorch import Callback, LightningModule, Trainer | |
| from lightning.pytorch.cli import LightningArgumentParser | |
| from lightning.pytorch.loggers import WandbLogger | |
| from smirk import SmirkTokenizerFast | |
| from torch import nn | |
| from torch.masked import MaskedTensor, masked_tensor | |
| from torchmetrics import MetricCollection | |
| from transformers import (AutoConfig, AutoModel, AutoTokenizer, | |
| DataCollatorWithPadding, PretrainedConfig, | |
| PreTrainedModel) | |
| AutoTokenizer.register("SmirkTokenizer", fast_tokenizer_class=SmirkTokenizerFast) | |
| MODEL_TYPE_ALIASES = {} | |
| IGNORE_INDEX = -100 | |
| def build_encoder(enc_dict: Dict[str, Any]): | |
| mtype = enc_dict.get("model_type") | |
| if mtype: | |
| base = MODEL_TYPE_ALIASES.get(mtype, mtype) | |
| cfg_cls = AutoConfig.for_model(base) | |
| enc_cfg = cfg_cls.from_dict(enc_dict) | |
| elif enc_dict.get("_name_or_path"): | |
| enc_cfg = AutoConfig.from_pretrained(enc_dict["_name_or_path"]) | |
| else: | |
| raise KeyError("encoder config missing 'model_type' or '_name_or_path'") | |
| if hasattr(enc_cfg, "add_pooling_layer"): | |
| enc_cfg.add_pooling_layer = False | |
| return AutoModel.from_config(enc_cfg) | |
| class AbstractNormalizer(torch.nn.Module): | |
| def __init__(self, num_outputs=None): | |
| super().__init__() | |
| self.num_outputs = num_outputs | |
| def forward(self, x): | |
| """Remove normalization""" | |
| raise NotImplementedError | |
| def inverse(self, x): | |
| """Apply normalization""" | |
| raise NotImplementedError | |
| def _fit(self, x): | |
| """Fit the normalization parameters""" | |
| raise NotImplementedError | |
| def to_config(self): | |
| return {'class': self.__class__.__name__, 'num_outputs': self.num_outputs} | |
| def leader_fit(self, ds, rank, broadcast): | |
| state = None | |
| if rank == 0: | |
| state = self.fit(ds) | |
| state = broadcast(state) | |
| self.load_state_dict(state) | |
| def fit(self, ds, name='target'): | |
| """Fit the normalization parameters on dataset""" | |
| if isinstance(ds, IterableDataset): | |
| target = [] | |
| mask = [] | |
| for x in ds: | |
| target.append(x[name]) | |
| mask.append(x[f'{name}_mask']) | |
| target = torch.stack(target) | |
| mask = torch.stack(mask) | |
| else: | |
| target = torch.stack([torch.tensor(x) for x in ds[name]]) | |
| mask = torch.stack([torch.tensor(x) for x in ds[f'{name}_mask']]) | |
| target = masked_tensor(target, mask) | |
| state = self._fit(target) | |
| return state | |
| def get(cls, transform, num_outputs): | |
| if isinstance(transform, list): | |
| assert len(transform) == num_outputs | |
| return ChannelWiseTransform([cls.get(t, 1) for t in transform]) | |
| elif transform in ['standardize', Standardize.__name__]: | |
| return Standardize(num_outputs) | |
| elif transform in ['power_transform', PowerTransform.__name__]: | |
| return PowerTransform(num_outputs) | |
| elif transform in ['log_transform', LogTransform.__name__]: | |
| return LogTransform(num_outputs) | |
| elif transform in ['max_scale', MaxScaleTransform.__name__]: | |
| return MaxScaleTransform(num_outputs) | |
| else: | |
| return IdentityTransform() | |
| class ArrheniusTaskHead(nn.Module): | |
| def __init__(self, embed_dim): | |
| super().__init__() | |
| self.desc_skip_connection = True | |
| self.fc1 = nn.Linear(embed_dim, embed_dim) | |
| self.relu1 = nn.GELU() | |
| self.fc2 = nn.Linear(embed_dim, embed_dim) | |
| self.relu2 = nn.GELU() | |
| self.fc3 = nn.Linear(embed_dim, int(0.5 * embed_dim)) | |
| self.relu3 = nn.GELU() | |
| self.final = nn.Linear(int(0.5 * embed_dim), 2) | |
| def forward(self, emb, temperature): | |
| x_out = self.fc1(emb) | |
| x_out = self.relu1(x_out) | |
| if self.desc_skip_connection is True: | |
| x_out = x_out + emb | |
| z = self.fc2(x_out) | |
| z = self.relu2(z) | |
| z = self.fc3(z) | |
| z = self.relu3(z) | |
| z = self.final(z) | |
| logA = z[:, 0] | |
| Ea = z[:, 1] | |
| R = 8.63e-05 | |
| e = torch.exp(torch.tensor([1], device=logA.device)) | |
| C = torch.log10(e) / R | |
| cond = logA - C * Ea / temperature | |
| return cond | |
| class ArrtheniusActivation(nn.Module): | |
| def forward(self, x, temperature): | |
| (A, Ea) = torch.chunk(x, 2, -1) | |
| return A * torch.exp(Ea / temperature) | |
| class BiPairwiseBlock(nn.Module): | |
| def __init__(self, d_model, bias=True, device=None, dtype=None): | |
| super().__init__() | |
| factory_kwargs = {'device': device, 'dtype': dtype} | |
| self.bi_weight = nn.Parameter(torch.empty((d_model, d_model), **factory_kwargs)) | |
| self.lin_weight = nn.Parameter(torch.empty((d_model, d_model), **factory_kwargs)) | |
| if bias: | |
| self.bias = nn.Parameter(torch.empty(d_model, **factory_kwargs)) | |
| else: | |
| self.register_parameter('bias', None) | |
| self.reset_parameters() | |
| self.bi_weight.register_hook(lambda grad: 0.5 * (grad + grad.T)) | |
| def reset_parameters(self): | |
| nn.init.xavier_normal_(self.lin_weight, gain=nn.init.calculate_gain('relu')) | |
| nn.init.xavier_normal_(self.bi_weight, gain=nn.init.calculate_gain('relu')) | |
| with torch.no_grad(): | |
| self.bi_weight.copy_(0.5 * (self.bi_weight + self.bi_weight.T)) | |
| if self.bias is not None: | |
| bound = 1 / math.sqrt(self.bias.size(0)) | |
| nn.init.uniform_(self.bias, -bound, bound) | |
| def forward(self, x): | |
| y_bi = torch.einsum('...ld,df,...rf->...lrf', x, self.bi_weight, x) | |
| y_bi = 0.5 * (y_bi + y_bi.transpose(-3, -2)) | |
| x_linear = x.unsqueeze(-2) + x.unsqueeze(-3) | |
| return y_bi + F.linear(x_linear, self.lin_weight, self.bias) | |
| class CanSkip: | |
| def should_skip(self): | |
| """Return true if the model should skip this batch, the model | |
| should still perform a forward pass, but return `0 * loss` | |
| instead. Do not return `None` as unsupported by DeepSpeed | |
| """ | |
| if hasattr(self, 'skip_this_batch') and self.skip_this_batch: | |
| return True | |
| else: | |
| return False | |
| class ChannelWiseTransform(AbstractNormalizer): | |
| def __init__(self, transforms): | |
| super().__init__(len(transforms)) | |
| self.transforms = torch.nn.ModuleList(transforms) | |
| def to_config(self): | |
| return {'class': [t.__class__.__name__ for t in self.transforms], 'num_outputs': self.num_outputs} | |
| def inverse(self, x): | |
| return torch.cat([transform.inverse(x[:, [idx]]) for (idx, transform) in enumerate(self.transforms)], dim=1) | |
| def forward(self, x): | |
| return torch.cat([transform.forward(x[:, [idx]]) for (idx, transform) in enumerate(self.transforms)], dim=1) | |
| def _fit(self, x): | |
| for (idx, transform) in enumerate(self.transforms): | |
| transform._fit(x[:, [idx]]) | |
| return self.state_dict() | |
| class DeepSpeedMixin: | |
| def load(checkpoint_dir, **kwargs): | |
| print(checkpoint_dir) | |
| return SaveConfigWithCkpts.load(checkpoint_dir, **kwargs) | |
| def load_state(self, checkpoint_dir): | |
| print('Loading state for checkpoint:', checkpoint_dir) | |
| state = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) | |
| self.load_state_dict(state, strict=False, assign=True) | |
| def get_encoder(self): | |
| raise NotImplementedError | |
| class IdentityTransform(AbstractNormalizer): | |
| def inverse(self, x): | |
| return x | |
| def forward(self, x): | |
| return x | |
| def _fit(self, x): | |
| return self.state_dict() | |
| class LagrangePolynomial(nn.Module): | |
| def __init__(self, polynomial_order=4, zero_endpoints=False): | |
| super().__init__() | |
| self.polynomial_order = polynomial_order | |
| self.zero_endpoints = zero_endpoints | |
| assert self.polynomial_order >= (2 if not zero_endpoints else 3) | |
| nodes = torch.tensor([cos((2 * k + 1) * pi / (2 * self.polynomial_order)) for k in range(self.polynomial_order)]) | |
| nodes -= nodes.min() | |
| nodes /= nodes.max() | |
| self.register_buffer('nodes', nodes, persistent=False) | |
| self.register_buffer('weights', self.barycentric_weights(self.nodes), persistent=False) | |
| self.register_buffer('weights_mask', ~torch.eye(self.polynomial_order, dtype=torch.bool), persistent=False) | |
| def active_nodes(self): | |
| if self.zero_endpoints: | |
| return self.nodes[1:-1] | |
| return self.nodes | |
| def barycentric_weights(cls, nodes): | |
| """Compute barycentric weights for Lagrange interpolation""" | |
| assert nodes.ndim == 1 | |
| x_diff = nodes.unsqueeze(0) - nodes.unsqueeze(1) | |
| x_diff.fill_diagonal_(1.0) | |
| w = 1.0 / x_diff.prod(dim=1) | |
| return w | |
| def forward(self, coefs, x): | |
| if self.zero_endpoints: | |
| zeros = torch.zeros_like(coefs[..., :1]) | |
| coefs = torch.cat([zeros, coefs, zeros], dim=-1) | |
| assert coefs.shape[-1] == self.polynomial_order | |
| x = x.unsqueeze(-1) | |
| nodes = self.nodes | |
| n = self.polynomial_order | |
| x_diff = nodes - x | |
| x_diff_masked = x_diff.unsqueeze(-2).expand(*x.shape[:-1], n, n) | |
| x_diff_masked = x_diff_masked.masked_select(self.weights_mask) | |
| x_diff_masked = x_diff_masked.view(*x.shape[:-1], n, n - 1) | |
| num = x_diff_masked.prod(dim=-1) | |
| basis = num * self.weights | |
| y = (basis * coefs).sum(dim=-1) | |
| return y | |
| class LinearExogenousEffect(nn.Module): | |
| def forward(self, x, temperature): | |
| (a, b) = torch.chunk(x, 2, -1) | |
| return a + b * (temperature / 293.15) | |
| class LoggingMixin: | |
| def on_train_epoch_start(self): | |
| self.trainer.train_dataloader.dataset.set_epoch(self.trainer.current_epoch) | |
| self.log('train/dataloader_epoch', self.trainer.train_dataloader.dataset._epoch, rank_zero_only=True, sync_dist=True) | |
| return super().on_train_epoch_start() | |
| class MISTExcessPhysicsConfig(PretrainedConfig): | |
| model_type = 'mist_excess_physics' | |
| def __init__(self, encoder=None, interactions='difference', num_control=3, num_targets=1, temperature_dependence=None, relative_excess=False, dropout=0.1, tokenizer_class='SmirkTokenizer', **kwargs): | |
| super().__init__(**kwargs) | |
| self.encoder = encoder or {} | |
| self.interactions = interactions | |
| self.num_control = num_control | |
| self.num_targets = num_targets | |
| self.temperature_dependence = temperature_dependence | |
| self.relative_excess = relative_excess | |
| self.dropout = dropout | |
| self.tokenizer_class = tokenizer_class | |
| class MISTExcessPhysics(PreTrainedModel): | |
| config_class = MISTExcessPhysicsConfig | |
| def __init__(self, config): | |
| super().__init__(config) | |
| self.config = config | |
| self.encoder = build_encoder_from_dict(config.encoder) | |
| n_env = 0 | |
| n_temperature_targets = 1 | |
| if config.temperature_dependence == 'arrhenius': | |
| n_temperature_targets = 2 | |
| self.temperature_dependence = ArrtheniusActivation() | |
| elif config.temperature_dependence == 'concat': | |
| n_env = 1 | |
| self.temperature_dependence = lambda x, t: x | |
| elif config.temperature_dependence == 'locally-linear': | |
| n_temperature_targets = 2 | |
| n_env = 1 | |
| self.temperature_dependence = LinearExogenousEffect() | |
| else: | |
| self.temperature_dependence = lambda x, t: x | |
| self.pairwise_interaction = pairwise_fusion(config.interactions, config.encoder['hidden_size'], config.num_control * n_temperature_targets, n_targets=config.num_targets, n_env=n_env) | |
| self.component_properties = nn.Sequential(nn.Linear(self.encoder.config.hidden_size + n_env, self.encoder.config.hidden_size), nn.Dropout(config.dropout), nn.SiLU(), nn.Linear(self.encoder.config.hidden_size, config.num_targets * n_temperature_targets)) | |
| self.excess_polynomial = LagrangePolynomial(polynomial_order=config.num_control + 2, zero_endpoints=True) | |
| self.transform = Standardize(num_outputs=config.num_targets) | |
| self.excess_transform = Standardize(num_outputs=config.num_targets) | |
| self.tokenizer = None | |
| self.post_init() | |
| def from_components(cls, encoder, pairwise_interaction, component_properties, excess_polynomial, transform, excess_transform, tokenizer=None, interactions='difference', num_control=3, num_targets=1, temperature_dependence=None, relative_excess=False, dropout=0.1): | |
| cfg = MISTExcessPhysicsConfig(encoder=encoder.config.to_dict(), interactions=interactions, num_control=num_control, num_targets=num_targets, temperature_dependence=temperature_dependence, relative_excess=relative_excess, dropout=dropout, tokenizer_class=getattr(tokenizer, '__class__', type('T', (), {})).__name__ if tokenizer else 'SmirkTokenizer') | |
| model = cls(cfg) | |
| model.encoder.load_state_dict(encoder.state_dict(), strict=False) | |
| model.pairwise_interaction.load_state_dict(pairwise_interaction.state_dict()) | |
| model.component_properties.load_state_dict(component_properties.state_dict()) | |
| model.excess_polynomial.load_state_dict(excess_polynomial.state_dict()) | |
| model.transform.load_state_dict(transform.state_dict()) | |
| model.excess_transform.load_state_dict(excess_transform.state_dict()) | |
| model.tokenizer = tokenizer | |
| return model | |
| def forward(self, input_ids, attention_mask, composition, temperature): | |
| hs = self.encoder(input_ids.reshape(-1, input_ids.shape[-1]), attention_mask=attention_mask.reshape(-1, attention_mask.shape[-1]), return_dict=True, output_attentions=False).last_hidden_state.reshape(*input_ids.shape, -1) | |
| embs = masked_mean_pool(hs, attention_mask) | |
| (y, y_linear, y_excess) = self.compute_interactions(composition, embs, temperature) | |
| return (y, y_linear, y_excess) | |
| def compute_interactions(self, composition, embs, temperature): | |
| (B, C, E) = embs.shape | |
| indices = torch.triu_indices(C, C, offset=1) | |
| I_ = indices.shape[1] | |
| e_i = embs[:, indices[0]].reshape(B * I_, E) | |
| e_j = embs[:, indices[1]].reshape(B * I_, E) | |
| t_ij = temperature.view(B, 1, 1).expand(-1, I_, -1).reshape(B * I_, 1) | |
| pw_coeffs = self.pairwise_interaction(e_i, e_j, t_ij) | |
| x_t = composition[:, indices[0]] + composition[:, indices[1]] | |
| x_t = x_t.clamp(min=0, max=1) | |
| x_i = composition[:, indices[0]] / x_t | |
| x_i = torch.where(x_i.abs() > 1e-08, x_i, torch.tensor(0.0).to(x_i)) | |
| x_i = x_i.clamp(min=0, max=1) | |
| x_i = x_i.view(B * I_, 1) | |
| pw_coeffs = self.temperature_dependence(pw_coeffs, t_ij.view(B * I_, 1, 1)) | |
| pw = x_t.view(B * I_, 1) * self.excess_polynomial(pw_coeffs, x_i) | |
| y_excess = pw.reshape(B, I_, -1).sum(dim=1) | |
| if self.config.temperature_dependence in ['concat', 'locally-linear']: | |
| t_e = temperature.view(B, 1, 1).expand(-1, C, -1) | |
| y_target = self.component_properties(torch.cat([embs, t_e], dim=-1)) | |
| else: | |
| y_target = self.component_properties(embs) | |
| y_target = self.temperature_dependence(y_target, temperature.view(B, 1, 1)) | |
| y_linear = (y_target * composition.view(B, C, 1)).sum(dim=1) | |
| if self.config.relative_excess: | |
| y_excess *= y_linear | |
| y_linear = self.transform.forward(y_linear) | |
| y_excess = self.excess_transform.forward(y_excess) | |
| y = y_linear + y_excess | |
| return (y, y_linear, y_excess) | |
| def predict(self, smiles_list, composition, temperature): | |
| tok = resolve_tokenizer(self, None) | |
| all_smiles = [smi for mixture in smiles_list for smi in mixture] | |
| inputs = tok(all_smiles, padding='longest', return_tensors='pt') | |
| batch_size = len(smiles_list) | |
| n_components = len(smiles_list[0]) | |
| seq_len = inputs['input_ids'].shape[-1] | |
| input_ids = inputs['input_ids'].reshape(batch_size, n_components, seq_len) | |
| attention_mask = inputs['attention_mask'].reshape(batch_size, n_components, seq_len) | |
| composition_tensor = torch.tensor(composition, dtype=torch.float32) | |
| temperature_tensor = torch.tensor(temperature, dtype=torch.float32) | |
| device = next(self.parameters()).device | |
| input_ids = input_ids.to(device) | |
| attention_mask = attention_mask.to(device) | |
| composition_tensor = composition_tensor.to(device) | |
| temperature_tensor = temperature_tensor.to(device) | |
| with torch.no_grad(): | |
| (y, y_linear, y_excess) = self.forward(input_ids=input_ids, attention_mask=attention_mask, composition=composition_tensor, temperature=temperature_tensor) | |
| return {'value': y.cpu(), 'linear': y_linear.cpu(), 'excess': y_excess.cpu()} | |
| def save_pretrained(self, save_directory, **kwargs): | |
| cfg = MISTExcessPhysicsConfig(encoder=self.encoder.config.to_dict(), interactions=self.config.interactions, num_control=self.config.num_control, num_targets=self.config.num_targets, temperature_dependence=self.config.temperature_dependence, relative_excess=self.config.relative_excess, dropout=self.config.dropout, tokenizer_class=self.tokenizer.__class__.__name__ if getattr(self, 'tokenizer', None) else 'SmirkTokenizer') | |
| super().save_pretrained(save_directory, config=cfg, **kwargs) | |
| if getattr(self, 'tokenizer', None) is not None: | |
| self.tokenizer.save_pretrained(save_directory) | |
| class MISTFinetunedConfig(PretrainedConfig): | |
| """HF config for a single-task MIST wrapper.""" | |
| model_type = 'mist_finetuned' | |
| def __init__(self, encoder=None, task_network=None, transform=None, channels=None, tokenizer_class='SmirkTokenizer', **kwargs): | |
| super().__init__(**kwargs) | |
| self.encoder = encoder or {} | |
| self.task_network = task_network or {} | |
| self.transform = transform or {} | |
| self.channels = channels | |
| self.tokenizer_class = tokenizer_class | |
| class MISTFinetuned(PreTrainedModel): | |
| config_class = MISTFinetunedConfig | |
| def __init__(self, config): | |
| super().__init__(config) | |
| self.encoder = build_encoder_from_dict(config.encoder) | |
| tn = config.task_network | |
| self.task_network = PredictionTaskHead(embed_dim=tn['embed_dim'], output_size=tn['output_size'], dropout=tn['dropout']) | |
| self.transform = AbstractNormalizer.get(config.transform['class'], config.transform['num_outputs']) | |
| self.channels = config.channels | |
| self.tokenizer = None | |
| self.post_init() | |
| def from_components(cls, encoder, task_network, transform, tokenizer=None, channels=None): | |
| cfg = MISTFinetunedConfig(encoder=encoder.config.to_dict(), task_network={'embed_dim': encoder.config.hidden_size, 'output_size': task_network.final.out_features, 'dropout': task_network.dropout1.p}, transform=transform.to_config(), channels=channels, tokenizer_class=getattr(tokenizer, '__class__', type('T', (), {})).__name__ if tokenizer else 'SmirkTokenizer') | |
| model = cls(cfg) | |
| model.encoder.load_state_dict(encoder.state_dict(), strict=False) | |
| model.task_network.load_state_dict(task_network.state_dict()) | |
| model.transform.load_state_dict(transform.state_dict()) | |
| model.tokenizer = tokenizer | |
| return model | |
| def forward(self, input_ids, attention_mask=None): | |
| hs = self.encoder(input_ids, attention_mask=attention_mask).last_hidden_state | |
| y = self.task_network(hs) | |
| return self.transform.forward(y) | |
| def _resolve_tokenizer(self, tokenizer): | |
| if tokenizer is not None: | |
| return tokenizer | |
| if getattr(self, 'tokenizer', None) is not None: | |
| return self.tokenizer | |
| try: | |
| return AutoTokenizer.from_pretrained(self.name_or_path, use_fast=True, trust_remote_code=True) | |
| except Exception: | |
| return AutoTokenizer.from_pretrained(self.config._name_or_path, use_fast=True, trust_remote_code=True) | |
| def embed(self, smi, tokenizer=None): | |
| tok = self._resolve_tokenizer(tokenizer) | |
| batch = tok(smi) | |
| batch = DataCollatorWithPadding(tok)(batch) | |
| input_ids = batch['input_ids'].to(self.device) | |
| attention_mask = batch['attention_mask'].to(self.device) | |
| with torch.inference_mode(): | |
| hs = self.encoder(input_ids, attention_mask=attention_mask).last_hidden_state[:, 0, :] | |
| return hs.to('cpu') | |
| def predict(self, smi, return_dict=True, tokenizer=None): | |
| tok = self._resolve_tokenizer(tokenizer) | |
| batch = tok(smi) | |
| batch = DataCollatorWithPadding(tok)(batch) | |
| inputs = {k: v.to(self.device) for (k, v) in batch.items()} | |
| with torch.inference_mode(): | |
| out = self(**inputs).cpu() | |
| if self.channels is None or not return_dict: | |
| return out | |
| return annotate_prediction(out, maybe_get_annotated_channels(self.channels)) | |
| def save_pretrained(self, save_directory, **kwargs): | |
| super().save_pretrained(save_directory, **kwargs) | |
| if getattr(self, 'tokenizer', None) is not None: | |
| self.tokenizer.save_pretrained(save_directory) | |
| class MISTMultiTaskConfig(PretrainedConfig): | |
| """HuggingFace config for a multi-task MIST wrapper.""" | |
| model_type = 'mist_multitask' | |
| def __init__(self, encoder=None, task_networks=None, transforms=None, channels=None, tokenizer_class='SmirkTokenizer', **kwargs): | |
| super().__init__(**kwargs) | |
| self.encoder = encoder or {} | |
| self.task_networks = task_networks or [] | |
| self.transforms = transforms or [] | |
| self.channels = channels | |
| self.tokenizer_class = tokenizer_class | |
| class MISTMultiTask(PreTrainedModel): | |
| config_class = MISTMultiTaskConfig | |
| def __init__(self, config): | |
| super().__init__(config) | |
| self.encoder = build_encoder_from_dict(config.encoder) | |
| self.task_networks = nn.ModuleList([PredictionTaskHead(embed_dim=tn['embed_dim'], output_size=tn['output_size'], dropout=tn['dropout']) for tn in config.task_networks]) | |
| self.transforms = nn.ModuleList([AbstractNormalizer.get(tf_cfg['class'], tf_cfg['num_outputs']) for tf_cfg in config.transforms]) | |
| assert len(self.task_networks) == len(self.transforms), 'task_networks and transforms must align' | |
| self.channels = config.channels | |
| self.tokenizer = None | |
| self.post_init() | |
| def from_components(cls, encoder, task_networks, transforms, tokenizer=None, channels=None): | |
| cfg = MISTMultiTaskConfig(encoder=encoder.config.to_dict(), task_networks=[{'embed_dim': encoder.config.hidden_size, 'output_size': tn.final.out_features, 'dropout': tn.dropout1.p} for tn in task_networks], transforms=[tf.to_config() for tf in transforms], channels=channels, tokenizer_class=getattr(tokenizer, '__class__', type('T', (), {})).__name__ if tokenizer else 'SmirkTokenizer') | |
| model = cls(cfg) | |
| model.encoder.load_state_dict(encoder.state_dict(), strict=False) | |
| for (dst, src) in zip(model.task_networks, task_networks): | |
| dst.load_state_dict(src.state_dict()) | |
| for (dst, src) in zip(model.transforms, transforms): | |
| dst.load_state_dict(src.state_dict()) | |
| model.tokenizer = tokenizer | |
| return model | |
| def forward(self, input_ids, attention_mask=None): | |
| hs = self.encoder(input_ids, attention_mask=attention_mask).last_hidden_state | |
| outs = [] | |
| for (tn, tf) in zip(self.task_networks, self.transforms): | |
| outs.append(tf.forward(tn(hs))) | |
| return torch.cat(outs, dim=-1) | |
| def _resolve_tokenizer(self, tokenizer): | |
| if tokenizer is not None: | |
| return tokenizer | |
| if getattr(self, 'tokenizer', None) is not None: | |
| return self.tokenizer | |
| try: | |
| return AutoTokenizer.from_pretrained(self.name_or_path, use_fast=True, trust_remote_code=True) | |
| except Exception: | |
| return AutoTokenizer.from_pretrained(self.config._name_or_path, use_fast=True, trust_remote_code=True) | |
| def predict(self, smi, tokenizer=None): | |
| tok = self._resolve_tokenizer(tokenizer) | |
| batch = tok(smi) | |
| batch = DataCollatorWithPadding(tok)(batch) | |
| inputs = {k: v.to(self.device) for (k, v) in batch.items()} | |
| with torch.inference_mode(): | |
| out = self(**inputs).cpu() | |
| if self.channels is None: | |
| return out | |
| return annotate_prediction(out, self.channels) | |
| def embed(self, smi, tokenizer=None): | |
| tok = self._resolve_tokenizer(tokenizer) | |
| batch = tok(smi) | |
| batch = DataCollatorWithPadding(tok)(batch) | |
| input_ids = batch['input_ids'].to(self.device) | |
| attention_mask = batch['attention_mask'].to(self.device) | |
| with torch.inference_mode(): | |
| hs = self.encoder(input_ids, attention_mask=attention_mask).last_hidden_state[:, 0, :] | |
| return hs.to('cpu') | |
| def save_pretrained(self, save_directory, **kwargs): | |
| super().save_pretrained(save_directory, **kwargs) | |
| if getattr(self, 'tokenizer', None) is not None: | |
| self.tokenizer.save_pretrained(save_directory) | |
| class MaxScaleTransform(AbstractNormalizer): | |
| """ | |
| Divide by maximum value in training dataset. | |
| """ | |
| def __init__(self, mx, eps=1e-08): | |
| super().__init__(1) | |
| self.num_outputs = 1 | |
| self.max = mx | |
| self.eps = float(eps) | |
| assert 0 <= self.eps | |
| def forward(self, x): | |
| x_out = self.max * x | |
| return x_out | |
| def inverse(self, x): | |
| x_out = x / self.max | |
| return x_out | |
| def _fit(self, target): | |
| return self.state_dict() | |
| class PairwiseInteraction(nn.Module): | |
| def __init__(self, n_in, n_out, n_targets=1, dropout=0.1, n_env=0): | |
| super().__init__() | |
| self.n_in = n_in | |
| self.n_out = n_out | |
| self.n_targets = n_targets | |
| self.n_env = n_env | |
| self.mlp_emb = nn.Sequential(nn.Linear(n_in, n_in), nn.Dropout(dropout)) | |
| self.mlp = nn.Sequential(nn.Linear(n_in + n_env, n_in), nn.Dropout(dropout), nn.GELU(), nn.Linear(n_in, n_out * n_targets)) | |
| self.reset_parameters() | |
| def reset_parameters(self): | |
| def init_weights(m): | |
| if isinstance(m, nn.Linear): | |
| nn.init.xavier_normal_(m.weight, gain=nn.init.calculate_gain('relu')) | |
| if m.bias is not None: | |
| nn.init.zeros_(m.bias) | |
| last = self.mlp[-1] | |
| nn.init.xavier_normal_(last.weight, gain=nn.init.calculate_gain('linear')) | |
| self.mlp.apply(init_weights) | |
| def forward(self, a, b, e=None): | |
| a = self.mlp_emb(a) | |
| b = self.mlp_emb(b) | |
| d = self.distance(a, b) | |
| if self.n_env > 0: | |
| d = torch.cat([d, e], dim=-1) | |
| y = self.mlp(d) | |
| y = y.reshape(*y.shape[:-1], self.n_targets, self.n_out) | |
| return y + y.flip(-1) | |
| def distance(self, a, b): | |
| return (a - b).pow(2) | |
| class EquivariantInteraction(PairwiseInteraction): | |
| def forward(self, a, b, e=None): | |
| a = self.mlp_emb(a) | |
| b = self.mlp_emb(b) | |
| (emb_a, emb_b) = self.distance(a, b) | |
| if self.n_env > 0: | |
| emb_a = torch.cat([emb_a, e], dim=-1) | |
| emb_b = torch.cat([emb_b, e], dim=-1) | |
| y_a = self.mlp(emb_a).reshape(*emb_a.shape[:-1], self.n_targets, self.n_out) | |
| y_b = self.mlp(emb_b).reshape(*emb_b.shape[:-1], self.n_targets, self.n_out) | |
| return y_a + y_b.flip(-1) | |
| def distance(self, a, b): | |
| return (a - b, b - a) | |
| class ConcatFusion(EquivariantInteraction): | |
| def __init__(self, n_in, n_out, **kwargs): | |
| super().__init__(n_in, n_out, **kwargs) | |
| self.mlp[0] = nn.Linear(2 * self.n_in + self.n_env, self.n_in) | |
| self.reset_parameters() | |
| def distance(self, a, b): | |
| return (torch.cat([a, b], dim=-1), torch.cat([b, a], dim=-1)) | |
| class GaussianFusion(PairwiseInteraction): | |
| def distance(self, a, b): | |
| dist = -1 * (a - b).pow(2) | |
| return dist.exp() | |
| class PairwiseMLP(nn.Module): | |
| def __init__(self, d_model, dropout=0.2, device=None, dtype=None): | |
| super().__init__() | |
| self.mlp = nn.Sequential(nn.Linear(2 * d_model, d_model), nn.Dropout(dropout), nn.GELU(), nn.Linear(d_model, d_model), nn.GELU()) | |
| def forward(self, x): | |
| (_, N, _) = x.shape | |
| x_l = x.unsqueeze(-2).expand(-1, N, N, -1) | |
| x_r = x.unsqueeze(-3).expand(-1, N, N, -1) | |
| x_pw = torch.cat([x_l, x_r], dim=-1) | |
| y = self.mlp(x_pw) | |
| return 0.5 * (y + y.transpose(1, 2)) | |
| class PolynomialPredictionTaskHead(nn.Module): | |
| def __init__(self, embed_dim, polynomial_order=4, dropout=0.1): | |
| super().__init__() | |
| self.polynomial_order = polynomial_order | |
| self.coeffients = nn.Sequential(nn.Linear(embed_dim, embed_dim), nn.Dropout(dropout), nn.GELU(), nn.Linear(embed_dim, polynomial_order)) | |
| def forward(self, emb, x): | |
| coefs = self.coeffients(emb) | |
| return self.eval_poly(coefs, x) | |
| def eval_poly(self, coefs, x): | |
| """Evaluate polynomial with coefficients `coefs` at `x`""" | |
| raise NotImplementedError | |
| class PowerTransform(AbstractNormalizer): | |
| """ | |
| Apply a power transform (Yeo-Johnson) featurewise to make data more Gaussian-like. | |
| Followed by applying a zero-mean, unit-variance normalization to the | |
| transformed output to rescale targets to [-1, 1]. | |
| """ | |
| def __init__(self, num_outputs, eps=1e-08): | |
| super().__init__(num_outputs) | |
| self.num_outputs = num_outputs | |
| self.register_buffer('lmbdas', torch.zeros(num_outputs)) | |
| self.register_buffer('mean', torch.zeros(num_outputs)) | |
| self.register_buffer('std', torch.zeros(num_outputs)) | |
| self.eps = float(eps) | |
| assert 0 <= self.eps | |
| def _yeo_johnson_transform(self, x, lmbda): | |
| """ | |
| Return transformed input x following Yeo-Johnson transform with | |
| parameter lambda. | |
| Adapted from | |
| https://github.com/scikit-learn/scikit-learn/blob/fbb32eae5/sklearn/preprocessing/_data.py#L3354 | |
| """ | |
| x_out = x.clone() | |
| eps = torch.finfo(x.dtype).eps | |
| pos = x >= 0 | |
| if abs(lmbda) < eps: | |
| x_out[pos] = torch.log1p(x[pos]) | |
| else: | |
| x_out[pos] = (torch.pow(x[pos] + 1, lmbda) - 1) / lmbda | |
| if abs(lmbda - 2) > eps: | |
| x_out[~pos] = -(torch.pow(-x[~pos] + 1, 2 - lmbda) - 1) / (2 - lmbda) | |
| else: | |
| x_out[~pos] = -torch.log1p(-x[~pos]) | |
| return x_out | |
| def _yeo_johnson_inverse_transform(self, x, lmbda): | |
| """ | |
| Return inverse-transformed input x following Yeo-Johnson inverse | |
| transform with parameter lambda. | |
| Adapted from | |
| https://github.com/scikit-learn/scikit-learn/blob/fbb32eae5/sklearn/preprocessing/_data.py#L3383 | |
| """ | |
| x_out = x.clone() | |
| pos = x >= 0 | |
| eps = torch.finfo(x.dtype).eps | |
| if abs(lmbda) < eps: | |
| x_out[pos] = torch.exp(x[pos]) - 1 | |
| else: | |
| x_out[pos] = torch.pow(x[pos] * lmbda + 1, 1 / lmbda) - 1 | |
| if abs(lmbda - 2) > eps: | |
| x_out[~pos] = 1 - torch.pow(-(2 - lmbda) * x[~pos] + 1, 1 / (2 - lmbda)) | |
| else: | |
| x_out[~pos] = 1 - torch.exp(-x[~pos]) | |
| return x_out | |
| def forward(self, x): | |
| x = self.std * x + self.mean | |
| x_out = torch.zeros_like(x) | |
| for i in range(self.num_outputs): | |
| x_out[:, i] = self._yeo_johnson_inverse_transform(x[:, i], self.lmbdas[i]) | |
| return x_out | |
| def inverse(self, x): | |
| x_out = torch.zeros_like(x) | |
| for i in range(self.num_outputs): | |
| x_out[:, i] = self._yeo_johnson_transform(x[:, i], self.lmbdas[i]) | |
| x_out = (x_out - self.mean) / self.std | |
| return x_out | |
| def _fit(self, target): | |
| from sklearn.preprocessing import PowerTransformer as _PowerTransformer | |
| transformer = _PowerTransformer(method='yeo-johnson', standardize=False) | |
| target = torch.tensor(transformer.fit_transform(target.get_data().numpy())) | |
| self.lmbdas = torch.tensor(transformer.lambdas_) | |
| self.mean = target.mean(0).to(self.mean) | |
| self.std = target.std(0).to(self.std) + self.eps | |
| return self.state_dict() | |
| class PredictionTaskHead(nn.Module): | |
| def __init__(self, embed_dim, output_size=1, dropout=0.2): | |
| super().__init__() | |
| self.desc_skip_connection = True | |
| self.fc1 = nn.Linear(embed_dim, embed_dim) | |
| self.dropout1 = nn.Dropout(dropout) | |
| self.relu1 = nn.GELU() | |
| self.fc2 = nn.Linear(embed_dim, embed_dim) | |
| self.dropout2 = nn.Dropout(dropout) | |
| self.relu2 = nn.GELU() | |
| self.final = nn.Linear(embed_dim, output_size) | |
| def forward(self, emb): | |
| if emb.ndim > 2: | |
| emb = emb[:, 0, :] | |
| x_out = self.fc1(emb) | |
| x_out = self.dropout1(x_out) | |
| x_out = self.relu1(x_out) | |
| if self.desc_skip_connection is True: | |
| x_out = x_out + emb | |
| z = self.fc2(x_out) | |
| z = self.dropout2(z) | |
| z = self.relu2(z) | |
| if self.desc_skip_connection is True: | |
| z = self.final(z + x_out) | |
| else: | |
| z = self.final(z) | |
| return z | |
| class SaveConfigWithCkpts(Callback): | |
| """Save Configuration with the model's checkpoints | |
| # Versions: Use Semantic Versioning for Model Checkpoint format | |
| ## Unnamed: | |
| - Stored `trainer.lightning_model.hparams` in model_hparms.json | |
| ## 0.2.0 | |
| - Added version field to model_hparams.json | |
| - Added "class_path" field | |
| - Moved model hparams to "init_args" field | |
| ## 0.2.1 | |
| - Save `JOB_CONFIG` to `job_config.json` | |
| ## 0.3.0 | |
| - Save hyperparameters under `lighting_module` and `datamodule`. | |
| """ | |
| VERSION = '0.3.0' | |
| def __init__(self, parser, config, overwrite=True): | |
| self.parser = parser | |
| self.config = config | |
| self.overwrite = overwrite | |
| self.already_saved = False | |
| self.config_path = None | |
| def setup(self, trainer, pl_module, stage): | |
| if self.already_saved: | |
| return | |
| self.config_path = self.log_dir(trainer) | |
| if trainer.is_global_zero: | |
| self.config_path.mkdir(parents=True, exist_ok=True) | |
| config_json = self.parser.dump(self.config, skip_none=False, skip_check=True, skip_link_targets=False, format='json') | |
| with open(Path(self.config_path, 'config.json'), 'w') as config_file: | |
| config_file.write(config_json) | |
| job_config = json.loads(os.environ.get('JOB_CONFIG', '{}')) | |
| with open(Path(self.config_path, 'job_config.json'), 'w') as fid: | |
| json.dump(job_config, fid) | |
| with open(Path(self.config_path, 'model_hparams.json'), 'w') as fid: | |
| model_cls = trainer.lightning_module.__class__ | |
| model_config = {'version': self.VERSION, 'class_path': f'{model_cls.__module__}.{model_cls.__name__}', 'lightning_module': trainer.lightning_module.hparams, 'datamodule': trainer.datamodule.hparams} | |
| json.dump(model_config, fid, default=lambda x: str(type(x))) | |
| with open(Path(self.config_path, 'env.json'), 'w') as fid: | |
| json.dump(dict(os.environ), fid, sort_keys=True) | |
| if (logger := trainer.logger): | |
| logger.log_hyperparams({'cli': self.config.as_dict()}) | |
| def log_dir(trainer): | |
| log_dir = trainer.log_dir or trainer.default_root_dir | |
| logger = trainer.logger | |
| if logger is not None and isinstance(logger, WandbLogger): | |
| config_path = Path(log_dir, str(logger.name), str(logger.version)) | |
| else: | |
| config_path = Path(log_dir) | |
| return config_path | |
| def instantiate(config_path, max_position_embeddings=None): | |
| """Instantiate a model from a checkpoint but don't load weights""" | |
| with open(config_path, 'r') as fid: | |
| config = json.load(fid) | |
| if (version := config.get('version', None)): | |
| if version.startswith('0.3'): | |
| (cls_name, model_config) = norm_class_config(config['lightning_module'], class_path=config.get('class_path', None)) | |
| model_config['vocab_size'] = config['datamodule']['vocab_size'] | |
| else: | |
| (cls_name, model_config) = norm_class_config(config) | |
| else: | |
| cls_name = 'electrolyte_fm.models.roberta_base.RoBERTa' | |
| model_config = config | |
| import_path = cls_name.split('.') | |
| if max_position_embeddings is not None: | |
| model_config['max_position_embeddings'] = max_position_embeddings | |
| model_cls = importlib.import_module('.'.join(import_path[:-1])).__getattribute__(import_path[-1]) | |
| assert import_path[-1] == model_cls.__name__ | |
| model = model_cls(**model_config) | |
| if hasattr(model, 'configure_model'): | |
| model.configure_model() | |
| return model | |
| def load(checkpoint_dir, config_path=None, map_location=None, strict=True, max_position_embeddings=None): | |
| """Restore from a deepspeed checkpoint, mainly used for downstream tasks""" | |
| checkpoint_dir = Path(checkpoint_dir).resolve() | |
| config_path = config_path or checkpoint_dir.parent.parent.joinpath('model_hparams.json') | |
| assert checkpoint_dir.exists(), f'Missing deepspeed checkpoint directory: {checkpoint_dir}' | |
| assert config_path.is_file(), f'Missing model config file {config_path}' | |
| model = SaveConfigWithCkpts.instantiate(config_path, max_position_embeddings) | |
| if not torch.cuda.is_available() and map_location is None: | |
| map_location = torch.device('cpu') | |
| if checkpoint_dir.is_file(): | |
| state = torch.load(checkpoint_dir, map_location=map_location, weights_only=False) | |
| if max_position_embeddings is not None: | |
| state = adjust_state_position_embeddings(state, max_position_embeddings) | |
| model.load_state_dict(state['state_dict'], strict=strict, assign=True) | |
| return model | |
| if checkpoint_dir.is_file(): | |
| state = torch.load(checkpoint_dir) | |
| model.load_state_dict(state['state_dict'], strict=strict, assign=True) | |
| return model | |
| try: | |
| from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint | |
| state = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) | |
| model.load_state_dict(state, strict=strict, assign=True) | |
| except FileNotFoundError: | |
| logging.error('failed to load checkpoint %s, trying to load rank 0 model states', checkpoint_dir) | |
| file = Path(checkpoint_dir, 'checkpoint', 'mp_rank_00_model_states.pt') | |
| state = torch.load(file) | |
| logging.info('loaded %s', file) | |
| model.load_state_dict(state['module'], strict=strict, assign=True) | |
| return model | |
| class SoftmaxFusion(EquivariantInteraction): | |
| def distance(self, a, b): | |
| y = torch.stack([a, b]).log_softmax(dim=0) | |
| return (y[0], y[1]) | |
| class Standardize(AbstractNormalizer): | |
| def __init__(self, num_outputs, eps=1e-08): | |
| super().__init__(num_outputs) | |
| self.register_buffer('mean', torch.zeros(num_outputs)) | |
| self.register_buffer('std', torch.zeros(num_outputs)) | |
| self.eps = float(eps) | |
| assert 0 <= self.eps | |
| def forward(self, x): | |
| return self.std * x + self.mean | |
| def inverse(self, x): | |
| return (x - self.mean) / self.std | |
| def fit(self, ds, name='target'): | |
| num_outputs = self.num_outputs | |
| assert num_outputs is not None | |
| mean = torch.zeros(num_outputs) | |
| m2 = torch.zeros(num_outputs) | |
| n = torch.zeros(num_outputs, dtype=torch.int) | |
| for row in ds: | |
| target = torch.tensor(row[name]) | |
| mask = torch.tensor(row[f'{name}_mask']) | |
| x = masked_tensor(target, mask) | |
| n += mask.view(-1, num_outputs).sum(0) | |
| xs = x.view(-1, num_outputs).sum(0) | |
| delta = xs - mean | |
| mean += (delta / n).get_data().masked_fill(~delta.get_mask(), 0) | |
| delta2 = xs - mean | |
| m2 += (delta * delta2).get_data().masked_fill(~delta.get_mask(), 0) | |
| self.mean = mean.to(self.mean) | |
| self.std = (m2 / n).sqrt().to(self.std) + self.eps | |
| self.mean[self.mean.isnan()] = 0 | |
| self.std[self.std.isnan()] = 1 | |
| logging.debug('Fitted %s', self.state_dict()) | |
| return self.state_dict() | |
| def _fit(self, target): | |
| self.mean = target.mean(0).get_data().to(self.mean) | |
| self.std = target.std(0).get_data().to(self.std) + self.eps | |
| return self.state_dict() | |
| def load_state_dict(self, state_dict, strict=True, assign=False): | |
| if 'transform.mean' in state_dict: | |
| state_dict = state_dict.copy() | |
| state_dict['mean'] = state_dict.pop('transform.mean') | |
| state_dict['std'] = state_dict.pop('transform.std') | |
| if assign: | |
| for (key, value) in state_dict.items(): | |
| if key in ['mean', 'std']: | |
| self.register_buffer(key, value) | |
| result = None | |
| else: | |
| result = super().load_state_dict(state_dict, strict=strict, assign=False) | |
| logging.debug(f' After loading: mean={self.mean}, std={self.std}') | |
| return result | |
| class LogTransform(Standardize): | |
| def forward(self, x): | |
| return torch.exp(super().forward(x)) | |
| def inverse(self, x): | |
| return super().inverse(torch.log(x)) | |
| def _fit(self, target): | |
| return super()._fit(torch.log(target)) | |
| class TokenPairwiseDistance(nn.Module): | |
| def __init__(self, embed_dim, dropout=0.2, num_attention_heads=1, num_layers=1, activation='relu', ff_ratio=2): | |
| super().__init__() | |
| enc_layer = nn.TransformerEncoderLayer(d_model=embed_dim, nhead=num_attention_heads, dim_feedforward=ff_ratio * embed_dim, dropout=dropout, batch_first=True, norm_first=True) | |
| self.interaction = nn.TransformerEncoder(enc_layer, num_layers) | |
| self.pairwise_distance = PairwiseMLP(embed_dim, dropout) | |
| self.distance1 = nn.Sequential(nn.Linear(embed_dim, embed_dim), nn.Dropout(dropout), nn.GELU()) | |
| self.distance2 = nn.Linear(embed_dim, 1) | |
| def forward(self, hs): | |
| hs = self.interaction(hs) | |
| with torch.autocast('cuda', dtype=torch.float32): | |
| pw_dist = self.pairwise_distance(hs) | |
| d = self.distance1(pw_dist) + pw_dist | |
| d = self.distance2(d).squeeze(-1) | |
| return F.relu(F.elu(d) + 1) | |
| class TokenTaskHead(nn.Module): | |
| def __init__(self, embed_dim, output_size=1, dropout=0.2): | |
| super().__init__() | |
| self.layers = nn.Sequential(nn.Linear(embed_dim, embed_dim), nn.Dropout(dropout), nn.GELU(), nn.Linear(embed_dim, embed_dim), nn.Dropout(dropout), nn.GELU(), nn.Linear(embed_dim, output_size)) | |
| def forward(self, emb): | |
| return self.layers(emb) | |
| class VFTDecayTaskHead(nn.Module): | |
| def __init__(self, embed_dim): | |
| super().__init__() | |
| self.desc_skip_connection = True | |
| self.fc1 = nn.Linear(embed_dim, embed_dim) | |
| self.relu1 = nn.GELU() | |
| self.fc2 = nn.Linear(embed_dim, embed_dim) | |
| self.relu2 = nn.GELU() | |
| self.fc3 = nn.Linear(embed_dim, int(0.5 * embed_dim)) | |
| self.relu3 = nn.GELU() | |
| self.final = nn.Linear(int(0.5 * embed_dim), 6) | |
| self.sigmoid = nn.Sigmoid() | |
| def forward(self, emb, temperature): | |
| x_out = self.fc1(emb) | |
| x_out = self.relu1(x_out) | |
| if self.desc_skip_connection is True: | |
| x_out = x_out + emb | |
| z = self.fc2(x_out) | |
| z = self.relu2(z) | |
| z = self.fc3(z) | |
| z = self.relu3(z) | |
| z = self.final(z) | |
| ln_A = z[:, 0] | |
| Ea = z[:, 1] | |
| Tg = z[:, 2] | |
| alpha = self.sigmoid(z[:, 3]) | |
| beta = z[:, 4] | |
| lmbda = self.sigmoid(z[:, 5]) | |
| params = {'conductivity': ln_A - Ea / (temperature - Tg), 'ln_A': ln_A, 'Ea': Ea, 'Tg': Tg, 'alpha': alpha, 'beta': beta, 'lmbda': lmbda} | |
| return params | |
| class VFTTaskHead(nn.Module): | |
| def __init__(self, embed_dim): | |
| super().__init__() | |
| self.desc_skip_connection = True | |
| self.fc1 = nn.Linear(embed_dim, embed_dim) | |
| self.relu1 = nn.GELU() | |
| self.fc2 = nn.Linear(embed_dim, embed_dim) | |
| self.relu2 = nn.GELU() | |
| self.fc3 = nn.Linear(embed_dim, int(0.5 * embed_dim)) | |
| self.relu3 = nn.GELU() | |
| self.final = nn.Linear(int(0.5 * embed_dim), 3) | |
| def forward(self, emb, temperature): | |
| x_out = self.fc1(emb) | |
| x_out = self.relu1(x_out) | |
| if self.desc_skip_connection is True: | |
| x_out = x_out + emb | |
| z = self.fc2(x_out) | |
| z = self.relu2(z) | |
| z = self.fc3(z) | |
| z = self.relu3(z) | |
| z = self.final(z) | |
| logA = z[:, 0] | |
| Ea = z[:, 1] | |
| T_g = z[:, 2] | |
| R = 8.63e-05 | |
| e = torch.exp(torch.tensor([1], device=logA.device)) | |
| C = torch.log10(e) / R | |
| cond = logA - C * Ea / (temperature - T_g) | |
| return cond | |
| def adjust_state_position_embeddings(state, max_position_embeddings): | |
| module_state = state['module'] | |
| weight_key = 'model.roberta_prelayernorm.embeddings.position_embeddings.weight' | |
| assert weight_key in module_state, 'Changing position embedding size only implemented for RoBERTaPreLayerNorm' | |
| (current_max_pos, embed_size) = module_state[weight_key].shape | |
| assert max_position_embeddings > current_max_pos, 'Maximum position embedding cannot be decreased' | |
| new_pos_embed = F.interpolate(module_state[weight_key].unsqueeze(0).permute(0, 2, 1), size=max_position_embeddings, mode='linear', align_corners=False).squeeze(0).permute(1, 0) | |
| module_state[weight_key] = new_pos_embed | |
| state['module'] = module_state | |
| return state | |
| def annotate_prediction(y, channels): | |
| out = {} | |
| for (idx, chn) in enumerate(channels): | |
| channel_info = {f: v for (f, v) in chn.items() if f != 'name'} | |
| out[chn['name']] = {'value': y[:, idx], **channel_info} | |
| return out | |
| def build_encoder_from_dict(enc_dict): | |
| if 'model_type' in enc_dict: | |
| cfg_cls = AutoConfig.for_model(enc_dict['model_type']) | |
| enc_cfg = cfg_cls.from_dict(enc_dict, strict=False) | |
| elif '_name_or_path' in enc_dict: | |
| enc_cfg = AutoConfig.from_pretrained(enc_dict['_name_or_path'], strict=False) | |
| else: | |
| raise KeyError("Encoder config is missing 'model_type' and '_name_or_path.") | |
| if hasattr(enc_cfg, 'add_pooling_layer'): | |
| enc_cfg.add_pooling_layer = False | |
| return AutoModel.from_config(enc_cfg) | |
| def get_ckpt_tokenizer(path): | |
| config_path = Path(path).parent.parent.joinpath('config.json') | |
| if not config_path.is_file(): | |
| return str(path) | |
| with open(config_path, 'r') as fid: | |
| config = json.load(fid) | |
| try: | |
| return config['data']['tokenizer'] | |
| except KeyError: | |
| return config['data']['init_args']['tokenizer'] | |
| def load_encoder(name_or_path, strict=False): | |
| if Path(name_or_path).exists(): | |
| return DeepSpeedMixin.load(name_or_path, strict=strict).get_encoder() | |
| else: | |
| from transformers import AutoModel | |
| return AutoModel.from_pretrained(name_or_path, trust_remote_code=True) | |
| def masked_mean_pool(x, mask): | |
| mask = mask.unsqueeze(-1) | |
| x_masked = x * mask | |
| s = x_masked.sum(dim=-2) | |
| c = mask.sum(dim=-2).clamp(min=1) | |
| return s / c | |
| def maybe_get_annotated_channels(channels): | |
| for chn in channels: | |
| if isinstance(chn, str): | |
| yield {'name': chn, 'description': None, 'unit': None} | |
| else: | |
| yield chn | |
| def norm_class_config(config, class_path=None): | |
| """Parse a dictionary of hparams for a class name and init args""" | |
| init_args = dict() | |
| if 'init_args' in config: | |
| init_args = config.pop('init_args') | |
| if 'class_path' in config: | |
| class_path = config['class_path'] | |
| elif '_class_path' in config: | |
| class_path = config.pop('_class_path') | |
| init_args = config | |
| init_args.pop('_instantiator', None) | |
| init_args.pop('instantiator', None) | |
| return (class_path, init_args) | |
| def pairwise_fusion(name, *args, **kwargs): | |
| if name == 'square-difference': | |
| return PairwiseInteraction(*args, **kwargs) | |
| elif name == 'gaussian': | |
| return GaussianFusion(*args, **kwargs) | |
| elif name == 'difference': | |
| return EquivariantInteraction(*args, **kwargs) | |
| elif name == 'softmax': | |
| return SoftmaxFusion(*args, **kwargs) | |
| elif name == 'concat': | |
| return ConcatFusion(*args, **kwargs) | |
| else: | |
| raise ValueError(f'Unknown fusion: {name}') | |
| def record_loss_summary_stats(logger): | |
| if isinstance(logger, WandbLogger): | |
| define_metric = logger.experiment.define_metric | |
| for m in ['train/loss', 'val/loss', 'test/loss']: | |
| for s in ['', '_step', '_epoch']: | |
| define_metric(m + s, summary='last,best,min', goal='minimize') | |
| def record_summary_stats(logger, metrics): | |
| if isinstance(logger, WandbLogger): | |
| define_metric = logger.experiment.define_metric | |
| for (name, metric) in metrics.items(): | |
| if not hasattr(metric, 'higher_is_better') or metric.higher_is_better is None: | |
| continue | |
| for phase in ['', '_epoch', '_step']: | |
| define_metric(name + phase, summary='last,best', goal='maximize' if metric.higher_is_better else 'minimize') | |
| def resolve_tokenizer(self, tokenizer=None): | |
| if tokenizer is not None: | |
| return tokenizer | |
| if getattr(self, 'tokenizer', None) is not None: | |
| return self.tokenizer | |
| try: | |
| return AutoTokenizer.from_pretrained(self.name_or_path, use_fast=True, trust_remote_code=True) | |
| except Exception: | |
| return AutoTokenizer.from_pretrained(self.config._name_or_path, use_fast=True, trust_remote_code=True) | |
| def sparsity_weights(ds, column, eps=1e-06, max=None): | |
| """Compute the relative sparsity of each channel in the target columns""" | |
| sparsity = {} | |
| for row in ds: | |
| for col in column: | |
| if col not in sparsity: | |
| sparsity[col] = torch.zeros_like(row[col], dtype=torch.float32) | |
| sparsity[col] += row[col] | |
| for col in sparsity.keys(): | |
| sparsity[col] = (sparsity[col].max() / (sparsity[col] + eps)).clamp(min=1, max=None) | |
| assert sparsity[col].min() >= 1 | |
| return sparsity | |
| class MISTIonicConductivityConfig(PretrainedConfig): | |
| model_type = 'mist_ionic_conductivity' | |
| def __init__(self, encoder=None, task_network=None, n_components=38, tokenizer_class='SmirkTokenizer', **kwargs): | |
| super().__init__(**kwargs) | |
| self.encoder = encoder or {} | |
| self.task_network = task_network or {'type': 'VFTDecayTaskHead', 'kwargs': {}} | |
| self.n_components = int(n_components) | |
| self.tokenizer_class = tokenizer_class | |
| class MISTIonicConductivity(PreTrainedModel): | |
| config_class = MISTIonicConductivityConfig | |
| def __init__(self, config): | |
| super().__init__(config) | |
| self.encoder = build_encoder_from_dict(config.encoder) | |
| tn_type = (config.task_network or {}).get('type', 'VFTDecayTaskHead') | |
| tn_kwargs = (config.task_network or {}).get('kwargs', {}) | |
| if tn_type == 'VFTDecayTaskHead': | |
| self.task_network = VFTDecayTaskHead(**tn_kwargs) | |
| else: | |
| raise ValueError(f'Unknown task_network type: {tn_type}') | |
| self.n_components = int(config.n_components) | |
| self.tokenizer = None | |
| self.post_init() | |
| def from_components(cls, encoder, task_network, tokenizer=None, n_components=38): | |
| if isinstance(task_network, VFTDecayTaskHead): | |
| tn = {'type': 'VFTDecayTaskHead', 'kwargs': {'embed_dim': encoder.config.hidden_size}} | |
| else: | |
| raise ValueError('Unsupported task head for MISTIonicConductivity.from_components') | |
| cfg = MISTIonicConductivityConfig(encoder=encoder.config.to_dict(), task_network=tn, n_components=n_components, tokenizer_class=getattr(tokenizer, '__class__', type('T', (), {})).__name__ if tokenizer else 'SmirkTokenizer') | |
| model = cls(cfg) | |
| model.encoder.load_state_dict(encoder.state_dict(), strict=False) | |
| model.task_network.load_state_dict(task_network.state_dict()) | |
| model.tokenizer = tokenizer | |
| return model | |
| def forward(self, batch, return_all=False): | |
| mix_embedding = None | |
| for i in range(self.n_components): | |
| enc = self.encoder(batch[f'input_ids_{i}'], attention_mask=batch[f'attention_mask_{i}'], return_dict=True, output_hidden_states=True).last_hidden_state[:, 0, :] | |
| comp = batch[f'composition_{i}'].view(-1, 1) | |
| enc = enc * comp | |
| mix_embedding = enc if mix_embedding is None else mix_embedding + enc | |
| params = self.task_network(mix_embedding, batch['temperature']) | |
| pred_unscaled = params['conductivity'] | |
| alpha = params['alpha'] | |
| beta = params['beta'] | |
| lmbda = params['beta'] | |
| exponent = (-1.0 * alpha + batch['composition_4']) / lmbda | |
| pred_decay = (1 - beta) * torch.exp(exponent) + beta | |
| pred = pred_unscaled * pred_decay | |
| pred = torch.where(batch['composition_4'] > alpha, pred, pred_unscaled) | |
| if return_all: | |
| return (pred.view(-1, 1), params) | |
| return (pred.view(-1, 1), alpha) | |
| def predict(self, batch, return_dict=True): | |
| """ | |
| Predict ionic conductivity for a batch of samples. | |
| """ | |
| tokenizer = resolve_tokenizer(self) | |
| collate = DataCollatorWithPadding(tokenizer) | |
| all_input_ids = [[] for _ in range(5)] | |
| all_attention_masks = [[] for _ in range(5)] | |
| all_compositions = [[] for _ in range(5)] | |
| all_temperatures = [] | |
| for sample in batch: | |
| solvent_composition = sample['solvent_composition'] | |
| cation = sample['cation'] | |
| anion = sample['anion'] | |
| temperature = sample['temperature'] | |
| total_solvent_comp = sum([v for comp_dict in solvent_composition for (k, v) in comp_dict.items()]) | |
| assert total_solvent_comp < 1.0, f'Total solvent mole fractions must be less than 1, got {total_solvent_comp}' | |
| salt_comp = 1.0 - total_solvent_comp | |
| components = list(solvent_composition) | |
| while len(components) < 3: | |
| components.append({'[H]': 0.0}) | |
| components.append({cation: salt_comp}) | |
| components.append({anion: salt_comp}) | |
| assert len(components) == 5, f'Expected 5 components, got {len(components)}' | |
| for (i, component) in enumerate(components): | |
| smiles = list(component.keys())[0] | |
| composition = list(component.values())[0] | |
| tok_output = tokenizer(smiles) | |
| all_input_ids[i].append(tok_output['input_ids']) | |
| all_attention_masks[i].append(tok_output['attention_mask']) | |
| all_compositions[i].append(composition) | |
| all_temperatures.append(temperature) | |
| output = {'temperature': torch.tensor(all_temperatures, dtype=torch.float32, device=self.device)} | |
| for i in range(5): | |
| batched = collate([{'input_ids': ids, 'attention_mask': mask} for (ids, mask) in zip(all_input_ids[i], all_attention_masks[i])]) | |
| output[f'input_ids_{i}'] = batched['input_ids'].to(self.device) | |
| output[f'attention_mask_{i}'] = batched['attention_mask'].to(self.device) | |
| output[f'composition_{i}'] = torch.tensor(all_compositions[i], dtype=torch.float32, device=self.device) | |
| with torch.inference_mode(): | |
| (pred, params) = self(output, return_all=return_dict) | |
| if not return_dict: | |
| return pred.cpu() | |
| return {'ln conductivity [mS/cm]': pred.cpu(), 'Ea': params['Ea'].cpu(), 'Tg': params['Tg'].cpu()} | |
| def save_pretrained(self, save_directory, **kwargs): | |
| tn = {'type': 'VFTDecayTaskHead', 'kwargs': {'embed_dim': self.encoder.config.hidden_size}} | |
| cfg = MISTIonicConductivityConfig(encoder=self.encoder.config.to_dict(), task_network=tn, n_components=self.n_components, tokenizer_class=self.tokenizer.__class__.__name__ if getattr(self, 'tokenizer', None) else 'SmirkTokenizer') | |
| super().save_pretrained(save_directory, config=cfg, **kwargs) | |
| if getattr(self, 'tokenizer', None) is not None: | |
| self.tokenizer.save_pretrained(save_directory) |