| """Pytorch Lightning implementation for VanillaTransformer.""" |
|
|
| import math |
| import json |
| import Levenshtein |
| import lightning as L |
| import torch |
| import torch.nn.functional as F |
| import torch.optim as optim |
| from typing import Dict |
| from torch import Tensor |
| from rdkit import Chem, DataStructs, RDLogger |
| from rdkit.Chem.rdFingerprintGenerator import GetMorganGenerator |
| from functools import lru_cache |
| from transformers import get_cosine_schedule_with_warmup |
| from types import MethodType |
| from transformers_model.smiles_tokenizer import SmilesTokenizer |
| from transformers_model.configuration import VanillaTransformerConfig |
| from transformers_model.transformer import VanillaTransformer |
| from transformers_model.logging_utils import setup_logger |
|
|
| logger = setup_logger(__name__, logging_level="info") |
| RDLogger.DisableLog('rdApp.*') |
|
|
|
|
| class LitVanillaTransformer(L.LightningModule): |
| """Pytorch lightning model for VanillaTransformer.""" |
|
|
| def __init__( |
| self, |
| vocab_path: str, |
| task: str, |
| learning_rate: float = 0.0001, |
| adam_beta1: float = 0.9, |
| adam_beta2: float = 0.98, |
| adam_epsilon: float = 1e-9, |
| adam_weight_decay: float = 0.01, |
| embedding_dim: int = 256, |
| ffnn_hidden_dim: int = 2048, |
| dropout: float = 0.1, |
| activation: str = "relu", |
| num_attention_heads: int = 8, |
| num_encoder_layers: int = 4, |
| num_decoder_layers: int = 4, |
| attention_mask: float = float("-inf"), |
| init_std: float = 0.02, |
| max_position_embeddings: int = 5000, |
| num_beams: int = 1, |
| topn: int = 1, |
| max_length: int = 256, |
| warmup_ratio: float = 0.06, |
| predictions_path: str = "data/predictions.jsonl", |
| ignore_nan: bool = False, |
| device: str = "cuda" |
| |
| ) -> None: |
| """Construct an LM lightning module. |
| """ |
| super().__init__() |
|
|
| self.save_hyperparameters() |
|
|
| self.model: VanillaTransformer |
| self.tokenizer = SmilesTokenizer(vocab_path) |
| self.task = task |
| self.learning_rate = learning_rate |
| self.adam_beta1 = adam_beta1 |
| self.adam_beta2 = adam_beta2 |
| self.adam_epsilon = adam_epsilon |
| self.adam_weight_decay = adam_weight_decay |
| self.embedding_dim = embedding_dim |
| self.ffnn_hidden_dim = ffnn_hidden_dim |
| self.dropout = dropout |
| self.activation = activation |
| self.num_attention_heads = num_attention_heads |
| self.num_encoder_layers = num_encoder_layers |
| self.num_decoder_layers = num_decoder_layers |
| self.attention_mask = attention_mask |
| self.init_std = init_std |
| self.max_position_embeddings = max_position_embeddings |
| self.num_beams = num_beams |
| self.topn = max(self.num_beams, topn) |
| self.max_length = max_length |
| self.warmup_ratio = warmup_ratio |
| self.predictions_path = predictions_path |
| self.ignore_nan = ignore_nan |
| self.device_pref = device |
| self.invalid_pred = 0 |
|
|
| self.init_model() |
|
|
| def init_model(self) -> None: |
| """Initialize a VanillaTransformer.""" |
|
|
| config_args = { |
| "pad_token_id": self.tokenizer.pad_token_id, |
| "eos_token_id": self.tokenizer.sep_token_id, |
| "bos_token_id": self.tokenizer.cls_token_id, |
| "decoder_start_token_id": self.tokenizer.cls_token_id, |
| "vocabulary_size": self.tokenizer.vocab_size, |
| "vocab_size": self.tokenizer.vocab_size, |
| "embedding_dim": self.embedding_dim, |
| "ffnn_hidden_dim": self.ffnn_hidden_dim, |
| "dropout": self.dropout, |
| "activation": self.activation, |
| "num_attention_heads": self.num_attention_heads, |
| "num_encoder_layers": self.num_encoder_layers, |
| "num_decoder_layers": self.num_decoder_layers, |
| "attention_mask": self.attention_mask, |
| "init_std": self.init_std, |
| "max_position_embeddings": self.max_position_embeddings, |
| "num_beams": self.num_beams, |
| "device": self.device_pref, |
| } |
| config = VanillaTransformerConfig(**config_args) |
| self.model = VanillaTransformer(config) |
|
|
| def forward(self, x: Tensor) -> Tensor: |
| """Forwards through the model. |
| """ |
| return self.model(**kwargs) |
|
|
| @lru_cache() |
| def total_steps(self): |
| return len(self.trainer.datamodule.train_dataloader()) // self.trainer.accumulate_grad_batches * self.trainer.max_epochs |
|
|
| def configure_optimizers( |
| self, |
| ) -> Dict[str, object]: |
| """Create and return the optimizer. |
| |
| Returns: |
| output (dict of str: Any): |
| - optimizer: the optimizer used to update the parameter. |
| """ |
|
|
| |
| optimizer = optim.AdamW( |
| params=self.parameters(), |
| lr=self.learning_rate, |
| betas=(self.adam_beta1, self.adam_beta2), |
| eps=self.adam_epsilon, |
| weight_decay=self.adam_weight_decay, |
| ) |
|
|
| total_steps = self.total_steps() |
| print("Total steps: ", total_steps) |
| warmup_steps = int(self.total_steps() * self.warmup_ratio) |
| print("Warmup steps: ", warmup_steps) |
|
|
| scheduler = get_cosine_schedule_with_warmup( |
| optimizer, |
| num_warmup_steps=warmup_steps, |
| num_training_steps=total_steps, |
| ) |
|
|
| output = { |
| "optimizer": optimizer, |
| "lr_scheduler": { |
| "scheduler": scheduler, |
| "interval": "step", |
| "frequency": 1 |
| } |
| } |
|
|
| return output |
|
|
| def training_step(self, batch: Dict[str, Tensor], batch_idx: int) -> Tensor: |
| """ |
| Training step which encompasses the forward pass and the computation of the loss value. |
| |
| Args: |
| batch: dictionary containing the input_ids and the attention_type. |
| batch_idx: index of the current batch, unused. |
| |
| Returns: |
| loss computed on the batch. |
| """ |
| loss = self.model(**batch).loss |
| self.log("train_loss", loss, on_step=False, on_epoch=True, logger=True) |
|
|
| |
| current_lr = self.trainer.optimizers[0].param_groups[0]['lr'] |
| self.log("current_lr", current_lr, on_step=False, on_epoch=True, logger=True) |
| |
| |
| if not self.ignore_nan: |
| if torch.isnan(loss): |
| raise ValueError("Train loss is NaN") |
|
|
| return loss |
|
|
| def on_validation_epoch_start(self): |
| if self.task == "forward": |
| self.invalid_pred = 0 |
|
|
| def validation_step(self, batch: Dict[str, Tensor], batch_idx: int) -> Tensor: |
| """ |
| Validation step which encompasses the forward pass and the computation of the loss value. |
| |
| Args: |
| batch: dictionary containing the input_ids and the attention_type. |
| batch_idx: index of the current batch, unused. |
| |
| Returns: |
| loss computed on the batch. |
| """ |
|
|
| |
| outputs = self.model(**batch) |
| loss = outputs.loss |
| self.log("val_loss", loss, on_step=False, on_epoch=True, logger=True) |
| |
| |
| if not self.ignore_nan: |
| if torch.isnan(loss): |
| raise ValueError("Validation loss is NaN") |
|
|
|
|
| sources = batch["encoder_input_ids"] |
| targets = batch["decoder_input_ids"] |
|
|
| |
| predictions = self.model.generate( |
| sources, |
| do_sample=False, |
| max_length=self.max_length, |
| num_beams=self.num_beams, |
| ) |
|
|
| |
| predicted_texts = self.tokenizer.batch_decode(predictions, skip_special_tokens=True) |
| predicted_texts = self.tokenizer.batch_detokenize_smiles_string(predicted_texts) |
| target_texts = self.tokenizer.batch_decode(targets, skip_special_tokens=True) |
| target_texts = self.tokenizer.batch_detokenize_smiles_string(target_texts) |
| |
| |
| correct = sum(pred == tgt for pred, tgt in zip(predicted_texts, target_texts)) |
| accuracy = correct / len(target_texts) |
| self.log("val_accuracy", accuracy, prog_bar=True, on_step=False, on_epoch=True, logger=True) |
|
|
| |
| levenshtein_distances = [ |
| Levenshtein.distance(pred, tgt) |
| for pred, tgt in zip(predicted_texts, target_texts) |
| ] |
| avg_levenshtein_distance = sum(levenshtein_distances) / len(levenshtein_distances) |
| self.log("val_levenshtein_distance_avg", avg_levenshtein_distance, on_step=False, on_epoch=True, logger=True) |
|
|
| |
| levenshtein_ratios = [ |
| Levenshtein.ratio(pred, tgt) |
| for pred, tgt in zip(predicted_texts, target_texts) |
| ] |
| avg_levenshtein_ratio = sum(levenshtein_ratios) / len(levenshtein_ratios) |
| self.log("val_levenshtein_ratio_avg", avg_levenshtein_ratio, on_step=False, on_epoch=True, logger=True) |
|
|
| if self.task == "forward": |
| similarities = [] |
| generator = GetMorganGenerator(radius=2, fpSize=2048) |
| for pred, tgt in zip(predicted_texts, target_texts): |
| mol1 = Chem.MolFromSmiles(pred) |
| mol2 = Chem.MolFromSmiles(tgt) |
| |
| if mol1 is None or mol2 is None: |
| if mol1 is None: |
| self.invalid_pred += 1 |
| continue |
|
|
| fp1 = generator.GetFingerprint(mol1) |
| fp2 = generator.GetFingerprint(mol2) |
| sim = DataStructs.TanimotoSimilarity(fp1, fp2) |
| similarities.append(sim) |
|
|
| avg_valid_tanimoto = sum(similarities) / len(similarities) if len(similarities) > 0 else 0 |
| avg_total_tanimoto = sum(similarities) / len(predicted_texts) |
| self.log("val_tanimoto_valid_avg", avg_valid_tanimoto, on_step=False, on_epoch=True, logger=True) |
| self.log("val_tanimoto_total_avg", avg_total_tanimoto, on_step=False, on_epoch=True, logger=True) |
| return {"val_loss": loss, "val_accuracy": accuracy, "val_levenshtein_distance_avg": avg_levenshtein_distance, "val_levenshtein_ratio_avg": avg_levenshtein_ratio, "val_tanimoto_valid_avg": avg_valid_tanimoto, "val_tanimoto_total_avg": avg_total_tanimoto} |
| else: |
| return {"val_loss": loss, "val_accuracy": accuracy, "val_levenshtein_distance_avg": avg_levenshtein_distance, "val_levenshtein_ratio_avg": avg_levenshtein_ratio} |
|
|
| def on_validation_epoch_end(self): |
| if self.task == "forward": |
| self.log("val_tanimoto_invalid_pred", self.invalid_pred) |
| self.invalid_pred = 0 |
|
|
| def test_step(self, batch: Dict[str, Tensor], batch_idx: int) -> float: |
| """ |
| Test step which encompasses the forward pass and the computation of the accuracy and the loss value. |
| |
| Args: |
| batch: dictionary containing the input_ids and the attention_type. |
| batch_idx: index of the current batch, unused. |
| |
| Returns: |
| accuracy computed on the batch. |
| """ |
|
|
| sources = batch["encoder_input_ids"] |
| targets = batch["decoder_input_ids"] |
|
|
| |
| predictions = self.model.generate( |
| sources, |
| do_sample=False, |
| max_length=self.max_length, |
| num_beams=self.num_beams, |
| ) |
|
|
| |
| predicted_texts = self.tokenizer.batch_decode(predictions, skip_special_tokens=True) |
| predicted_texts = self.tokenizer.batch_detokenize_smiles_string(predicted_texts) |
| target_texts = self.tokenizer.batch_decode(targets, skip_special_tokens=True) |
| target_texts = self.tokenizer.batch_detokenize_smiles_string(target_texts) |
| |
| |
| correct = sum(pred == tgt for pred, tgt in zip(predicted_texts, target_texts)) |
| accuracy = correct / len(target_texts) |
| self.log("accuracy", accuracy, prog_bar=True, on_step=False, on_epoch=True, logger=True) |
|
|
| |
| levenshtein_distances = [ |
| Levenshtein.distance(pred, tgt) |
| for pred, tgt in zip(predicted_texts, target_texts) |
| ] |
| avg_levenshtein_distance = sum(levenshtein_distances) / len(levenshtein_distances) |
| self.log("levenshtein_distance_avg", avg_levenshtein_distance, on_step=False, on_epoch=True, logger=True) |
|
|
| |
| levenshtein_ratios = [ |
| Levenshtein.ratio(pred, tgt) |
| for pred, tgt in zip(predicted_texts, target_texts) |
| ] |
| avg_levenshtein_ratio = sum(levenshtein_ratios) / len(levenshtein_ratios) |
| self.log("levenshtein_ratio_avg", avg_levenshtein_ratio, on_step=False, on_epoch=True, logger=True) |
|
|
| if self.task == "forward": |
| similarities = [] |
| generator = GetMorganGenerator(radius=2, fpSize=2048) |
| for pred, tgt in zip(predicted_texts, target_texts): |
| mol1 = Chem.MolFromSmiles(pred) |
| mol2 = Chem.MolFromSmiles(tgt) |
| |
| if mol1 is None or mol2 is None: |
| if mol1 is None: |
| self.invalid_pred += 1 |
| continue |
|
|
| fp1 = generator.GetFingerprint(mol1) |
| fp2 = generator.GetFingerprint(mol2) |
| sim = DataStructs.TanimotoSimilarity(fp1, fp2) |
| similarities.append(sim) |
|
|
| avg_valid_tanimoto = sum(similarities) / len(similarities) if len(similarities) > 0 else 0 |
| avg_total_tanimoto = sum(similarities) / len(predicted_texts) |
| self.log("tanimoto_valid_avg", avg_valid_tanimoto, on_step=False, on_epoch=True, logger=True) |
| self.log("tanimoto_total_avg", avg_total_tanimoto, on_step=False, on_epoch=True, logger=True) |
| return {"accuracy": accuracy, "levenshtein_distance_avg": avg_levenshtein_distance, "levenshtein_ratio_avg": avg_levenshtein_ratio, "tanimoto_valid_avg": avg_valid_tanimoto, "tanimoto_total_avg": avg_total_tanimoto} |
| else: |
| return {"accuracy": accuracy, "levenshtein_distance_avg": avg_levenshtein_distance, "levenshtein_ratio_avg": avg_levenshtein_ratio} |
|
|
| def on_predict_start(self): |
| |
| with open("predictions.jsonl", "w") as f: |
| f.truncate(0) |
|
|
| def predict_step(self, batch: Dict[str, Tensor], batch_idx: int) -> Dict: |
| """ |
| Predict step. |
| |
| Args: |
| batch: dictionary containing the input_ids and the attention_type. |
| batch_idx: index of the current batch, unused. |
| |
| Returns: |
| Predictions |
| """ |
|
|
| sources = batch["encoder_input_ids"] |
|
|
| |
| outputs = self.model.generate( |
| sources, |
| do_sample=False, |
| max_length=self.max_length, |
| num_beams=self.num_beams, |
| num_return_sequences=self.topn, |
| return_dict_in_generate=True, |
| output_scores=True, |
| ) |
|
|
| predictions = outputs.sequences |
| scores = outputs.sequences_scores |
|
|
| cross_mats_all = self._collect_cross_attentions(sources, predictions) |
|
|
| |
| predicted_texts = self.tokenizer.batch_decode(predictions, skip_special_tokens=True) |
| predicted_texts = self.tokenizer.batch_detokenize_smiles_string(predicted_texts) |
| target_texts = self.tokenizer.batch_decode(sources, skip_special_tokens=True) |
| target_texts = self.tokenizer.batch_detokenize_smiles_string(target_texts) |
|
|
| results = [] |
|
|
| |
| with open(self.predictions_path, "a") as f: |
| for i, target_text in enumerate(target_texts): |
| start = i * self.topn |
| end = start + self.topn |
| for j in range(start, end): |
| result = { |
| "source": target_text, |
| "predicted_target": predicted_texts[j], |
| "confidence": scores[j].item() |
| } |
|
|
| if cross_mats_all and j < len(cross_mats_all): |
| |
| result["cross_attentions"] = cross_mats_all[j] |
|
|
| |
| results.append(result) |
| |
| |
| return results |
|
|
| def _collect_cross_attentions(self, sources: torch.Tensor, predictions: torch.Tensor): |
| """ |
| Collect decoder→encoder cross-attention for each returned hypothesis. |
| Returns: List[List[torch.Tensor]] |
| out_all[b][l] is a tensor of shape [tgt_len, src_len] for hypothesis b and decoder layer l. |
| """ |
| model = self.model |
| device = sources.device |
|
|
| |
| enc_out = model.encode(input_ids=sources) |
| memory = enc_out.last_hidden_state |
|
|
| |
| B_eff = predictions.size(0) |
| B_mem = memory.size(0) |
| if B_mem != B_eff: |
| memory = memory.expand(B_eff, -1, -1) |
|
|
| |
| layers = list(model.decoder.decoder.layers) |
| orig_forwards = [] |
| for layer in layers: |
| mha = layer.multihead_attn |
| orig_forward = mha.forward |
|
|
| def forward_with_weights(self_mha, query, key, value, **kwargs): |
| kwargs['need_weights'] = True |
| kwargs.setdefault('average_attn_weights', True) |
| out = orig_forward(query, key, value, **kwargs) |
| layer._last_cross_attn = out[1] |
| return out |
|
|
| mha.forward = MethodType(forward_with_weights, mha) |
| orig_forwards.append(orig_forward) |
|
|
| try: |
| |
| _ = model.decoder( |
| input_ids=predictions, |
| encoder_output=memory, |
| padding_mask=None, |
| ) |
| finally: |
| |
| for layer, orig in zip(layers, orig_forwards): |
| layer.multihead_attn.forward = orig |
|
|
| |
|
|
| out_all = [] |
| num_layers = len(layers) |
|
|
| for b in range(B_eff): |
| last_mat = None |
| |
| for li in range(num_layers - 1, -1, -1): |
| attn = getattr(layers[li], "_last_cross_attn", None) |
| if attn is not None and attn.numel() > 0: |
| |
| last_mat = attn[b].detach().cpu() |
| break |
|
|
| if last_mat is None: |
| |
| last_mat = torch.empty(0, 0) |
|
|
| out_all.append(last_mat) |
|
|
| return out_all |
|
|