DualMedBert / modeling_dualmedbert.py
souvik-nlp's picture
Fix all_tied_weights_keys bug for newer transformers
3ad1e6d verified
Raw
History Blame Contribute Delete
3.75 kB
import torch
import torch.nn as nn
from transformers import PreTrainedModel, DistilBertConfig, DistilBertModel
from .configuration_dualmedbert import DualMedBertConfig
class LoRALinear(nn.Module):
def __init__(self, linear: nn.Linear, rank: int = 8, alpha: int = 32):
super().__init__()
self.linear = linear
self.linear.weight.requires_grad_(False)
if self.linear.bias is not None:
self.linear.bias.requires_grad_(False)
d_in = linear.in_features
d_out = linear.out_features
self.lora_A = nn.Parameter(torch.randn(d_in, rank) * 0.01)
self.lora_B = nn.Parameter(torch.zeros(rank, d_out))
self.scale = alpha / rank
def forward(self, x):
return self.linear(x) + (x @ self.lora_A @ self.lora_B) * self.scale
def inject_lora(model, rank: int = 8, alpha: int = 32):
for p in model.parameters():
p.requires_grad_(False)
for layer in getattr(model, "transformer").layer:
layer.attention.q_lin = LoRALinear(layer.attention.q_lin, rank, alpha)
layer.attention.v_lin = LoRALinear(layer.attention.v_lin, rank, alpha)
return model
class SelfAttnPool(nn.Module):
def __init__(self, hidden: int = 768, heads: int = 4):
super().__init__()
self.heads = heads
self.query = nn.Linear(hidden, heads, bias=False)
self.proj = nn.Linear(hidden * heads, hidden)
def forward(self, h: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
scores = self.query(h)
mask_3d = mask.unsqueeze(-1).expand_as(scores)
scores = scores.masked_fill(mask_3d == 0, torch.finfo(scores.dtype).min)
weights = torch.softmax(scores, dim=1)
contexts = torch.matmul(weights.transpose(1, 2), h)
flat = contexts.reshape(h.size(0), -1)
return self.proj(flat)
class DualMedBertForSequenceClassification(PreTrainedModel):
config_class = DualMedBertConfig
_tied_weights_keys = []
@property
def all_tied_weights_keys(self):
return {}
def __init__(self, config):
super().__init__(config)
self.num_classes = config.num_classes
base_config = DistilBertConfig()
base = DistilBertModel(base_config)
self.backbone = inject_lora(base, rank=config.lora_rank, alpha=config.lora_alpha)
self.attn_pool = SelfAttnPool(768, heads=4)
self.fuse = nn.Linear(768 * 2, 768)
self.relu = nn.ReLU()
self.fuse_drop = nn.Dropout(0.3)
self.dropout = nn.Dropout(0.2)
self.head = nn.Linear(768, self.num_classes)
def _backbone_pass(self, input_ids, attention_mask):
return self.backbone(input_ids=input_ids,
attention_mask=attention_mask).last_hidden_state
def forward(self, input_ids, attention_mask=None, **kwargs):
if attention_mask is None:
attention_mask = torch.ones_like(input_ids)
h = self._backbone_pass(input_ids, attention_mask)
cls = h[:, 0, :]
pool = self.attn_pool(h, attention_mask)
x = self.fuse(torch.cat([cls, pool], dim=-1))
x = self.relu(x)
x = self.fuse_drop(x)
x = self.dropout(x)
logits = self.head(x)
from transformers.modeling_outputs import SequenceClassifierOutput
return SequenceClassifierOutput(
loss=None,
logits=logits,
hidden_states=(h,),
attentions=None
)