Download module.py from Android12138/FlowResampler: direct link, hf CLI and curl.
- Browser
- Download file 8.53 kB
-
https://huggingface.co/Android12138/FlowResampler/resolve/main/module.py
- Command line
-
hf download hf://Android12138/FlowResampler/module.py
-
curl -L -o module.py https://huggingface.co/Android12138/FlowResampler/resolve/main/module.py
8.53 kB
| from dataclasses import asdict | |
| import torch | |
| from lightning.pytorch import LightningModule | |
| from torch_geometric.utils import dropout_edge | |
| from transformers import get_scheduler | |
| class BaseClassificationModule(LightningModule): | |
| """ | |
| Base PyTorch Lightning module to train a model for binary classification.. | |
| Handles shared operations for training, validation, and checkpointing. | |
| Args: | |
| config: The configuration object containing hyperparameters. | |
| model: The model to be trained. | |
| tokenizer: The tokenizer used for text processing. | |
| num_training_samples: The number of training samples in the dataset. Used for scheduler calculation. | |
| """ | |
| def __init__( | |
| self, | |
| config, | |
| model, | |
| tokenizer=None, | |
| num_training_samples=None, | |
| **kwargs, | |
| ): | |
| super().__init__() | |
| # Save the passed hyperparameters | |
| hparams = asdict(config) | |
| hparams.update(kwargs) # Save additional hyperparameters passed via kwargs | |
| hparams["num_training_samples"] = num_training_samples | |
| for k, v in hparams.items(): # Convert dtypes into string | |
| if isinstance(v, torch.dtype): | |
| hparams[k] = str(v) | |
| self.save_hyperparameters(hparams, ignore=["model", "tokenizer"]) | |
| # Model settings | |
| self.model = model | |
| self.tokenizer = tokenizer | |
| # Disable automatic optimization | |
| self.automatic_optimization = True | |
| def forward(self, *args, **kwargs): | |
| return self.model.forward(*args, **kwargs) | |
| def training_step(self, batch, batch_idx): | |
| # Forward pass | |
| outputs = self.forward( | |
| graphs=batch.get("graphs"), | |
| input_ids=batch.get("input_ids"), | |
| attention_mask=batch.get("attention_mask"), | |
| labels=batch.get("labels"), | |
| ) | |
| # Loss is computed inside the model when labels are provided (default is cross entropy loss) | |
| self.log("total_loss", outputs["loss"], prog_bar=True, on_step=True, on_epoch=False) | |
| return outputs["loss"] | |
| def validation_step(self, batch, batch_idx): | |
| # Forward pass (batch size is 1 during validation for stability) | |
| outputs = self.forward( | |
| graphs=batch.get("graphs"), | |
| input_ids=batch.get("input_ids"), | |
| attention_mask=batch.get("attention_mask"), | |
| labels=batch.get("labels"), | |
| ) | |
| # Logging (during evaluation, normally save by epoch rather than step) | |
| self.log("total_loss_val", outputs["loss"], prog_bar=True, on_step=False, on_epoch=True, batch_size=1) | |
| def configure_optimizers(self): | |
| # Set the log directory for saving predictions | |
| self.log_dir = self.trainer.log_dir | |
| optimizer = torch.optim.AdamW([p for p in self.model.parameters() if p.requires_grad], lr=self.hparams.lr) | |
| # Learning rate scheduler settings | |
| num_training_steps = self.hparams.num_training_samples * self.hparams.max_epochs // self.hparams.batch_size // self.hparams.accumulate_grad_batches | |
| num_warmup_steps = int(num_training_steps * self.hparams.warmup_ratio) | |
| lr_scheduler = get_scheduler( | |
| name=self.hparams.scheduler_name, | |
| optimizer=optimizer, | |
| num_warmup_steps=num_warmup_steps, | |
| num_training_steps=num_training_steps, | |
| ) | |
| return { | |
| "optimizer": optimizer, | |
| "lr_scheduler": { | |
| "scheduler": lr_scheduler, | |
| "interval": "step", | |
| "frequency": 1, | |
| }, | |
| } | |
| def on_load_checkpoint(self, checkpoint): | |
| # Manually load the state_dict into self.model | |
| state_dict = checkpoint["state_dict"] | |
| self.model.load_state_dict(state_dict, strict=False) | |
| def on_save_checkpoint(self, checkpoint): | |
| # Dynamically retrieve the names of all trainable parameters | |
| trainable_keys = { | |
| name for name, param in self.named_parameters() if param.requires_grad | |
| } | |
| # Filter the state_dict to keep only trainable weights and remove the 'model.' prefix | |
| checkpoint["state_dict"] = { | |
| (k[len("model."):] if k.startswith("model.") else k): v | |
| for k, v in checkpoint["state_dict"].items() | |
| if k in trainable_keys | |
| } | |
| class LMHeadConstrainModule(BaseClassificationModule): | |
| """ | |
| Binary classification via constrained decoding on the LM head. | |
| """ | |
| def __init__(self, *args, **kwargs): | |
| super().__init__(*args, **kwargs) | |
| # Specific model settings | |
| self.choice_ids = self.tokenizer(["false", "true"], add_special_tokens=False).input_ids | |
| def test_step(self, batch, batch_idx, dataloader_idx=0): | |
| # Forward pass (batch size is 1 during testing for stability) | |
| outputs = self.forward( | |
| graphs=batch.get("graphs"), | |
| input_ids=batch.get("prompts"), | |
| attention_mask=batch.get("prompt_attention_mask"), | |
| ) | |
| # Constrained decoding on choice_ids | |
| logits = outputs.logits[:, -1, self.choice_ids].squeeze(-1) # get next token logits, shape: (batch_size, 2) | |
| y_proba = torch.softmax(logits, dim=-1)[:, 1] # probability of "true", shape: (batch_size,) | |
| y_pred = (y_proba >= 0.5).int() # use threshold instead of torch.argmax(logits) | |
| # Return for the callback to save predictions and times | |
| return { | |
| "indices": batch.get("indices"), | |
| "y_pred": y_pred.tolist(), | |
| "y_proba": y_proba.tolist(), | |
| "logits": logits.tolist(), | |
| } | |
| class GraphRepresentationModule(BaseClassificationModule): | |
| """ | |
| Graph representation learning via a self-supervised objective. | |
| Extends BaseClassificationModule to include specific edge masking strategies for training and validation. | |
| """ | |
| def __init__(self, *args, **kwargs): | |
| super().__init__(*args, **kwargs) | |
| # Set dropout probability from hparams | |
| self.drop_edge_p = self.hparams.drop_edge_p | |
| def _prepare_masked_edges(self, graphs): | |
| """Helper method to apply dropout_edge.""" | |
| # Apply dropout to edges. Force training=True to generate the mask. | |
| masked_edge_index, retained_mask = dropout_edge( | |
| graphs.edge_index, | |
| p=self.drop_edge_p, | |
| force_undirected=False, | |
| training=True, | |
| ) | |
| # Directly filter the corresponding edge attributes using the boolean mask | |
| masked_edge_attr = graphs.edge_attr[retained_mask] | |
| # Identify the target edges for loss calculation (the dropped edges) | |
| dropped_mask = ~retained_mask | |
| target_edge_index = graphs.edge_index[:, dropped_mask] | |
| return masked_edge_index, masked_edge_attr, target_edge_index | |
| def training_step(self, batch, batch_idx): | |
| # Apply random edge masking for message passing to prevent data leakage | |
| masked_edge_index, masked_edge_attr, target_edge_index = self._prepare_masked_edges(batch.get("graphs")) | |
| # Forward pass | |
| outputs = self.forward( | |
| graphs=batch.get("graphs"), | |
| masked_edge_index=masked_edge_index, | |
| masked_edge_attr=masked_edge_attr, | |
| target_edge_index=target_edge_index, | |
| ) | |
| # Loss is computed inside the model | |
| self.log("total_loss", outputs["loss"], prog_bar=True, on_step=True, on_epoch=False, batch_size=batch.get("graphs").num_graphs) | |
| return outputs["loss"] | |
| def validation_step(self, batch, batch_idx): | |
| # Mask edges during validation as well to ensure the total_loss_val is valid. | |
| # The model needs to predict the full graphs.edge_index based on partial message passing. | |
| masked_edge_index, masked_edge_attr, target_edge_index = self._prepare_masked_edges(batch.get("graphs")) | |
| # Forward pass | |
| outputs = self.forward( | |
| graphs=batch.get("graphs"), | |
| masked_edge_index=masked_edge_index, | |
| masked_edge_attr=masked_edge_attr, | |
| target_edge_index=target_edge_index, | |
| ) | |
| # Logging (during evaluation, normally save by epoch rather than step) | |
| self.log("total_loss_val", outputs["loss"], prog_bar=True, on_step=False, on_epoch=True, batch_size=batch.get("graphs").num_graphs) | |