FlowResampler / module.py
Android12138's picture
Upload folder using huggingface_hub
1ac68b2 verified
Raw History Blame Contribute Delete
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)