import torch import torch.nn as nn from transformers import AutoModel class NeoAraBERTEntityPairTwoHead(nn.Module): # Entity-pair relation classifier with binary existence and 40-way positive heads. def __init__( self, model_name, num_positive_relations, no_relation_id, original_id_to_positive_id, type_dim=42, pair_hidden_1=1024, pair_hidden_2=512, dropout_p=0.1, existence_loss_weight=1.0, positive_relation_loss_weight=1.0, ): super().__init__() self.encoder = AutoModel.from_pretrained(model_name, trust_remote_code=True) hidden_size = self.encoder.config.hidden_size self.no_relation_id = int(no_relation_id) self.existence_loss_weight = float(existence_loss_weight) self.positive_relation_loss_weight = float(positive_relation_loss_weight) self.register_buffer( "original_id_to_positive_id", torch.tensor(original_id_to_positive_id, dtype=torch.long), ) self.pair_feature_dim = 5 * hidden_size + type_dim self.pair_mlp = nn.Sequential( nn.Linear(self.pair_feature_dim, pair_hidden_1), nn.LayerNorm(pair_hidden_1), nn.GELU(), nn.Dropout(dropout_p), nn.Linear(pair_hidden_1, pair_hidden_2), nn.LayerNorm(pair_hidden_2), nn.GELU(), nn.Dropout(dropout_p), ) self.existence_head = nn.Linear(pair_hidden_2, 2) self.positive_relation_head = nn.Linear(pair_hidden_2, num_positive_relations) def forward( self, input_ids, attention_mask, type_features, subject_marker_positions, object_marker_positions, labels=None, ): hidden = self.encoder( input_ids=input_ids, attention_mask=attention_mask, ).last_hidden_state batch_indices = torch.arange(hidden.shape[0], device=hidden.device) cls = hidden[:, 0, :] subject = hidden[batch_indices, subject_marker_positions.long(), :] object_ = hidden[batch_indices, object_marker_positions.long(), :] pair_features = torch.cat([ cls, subject, object_, torch.abs(subject - object_), subject * object_, type_features.float(), ], dim=-1) pair_hidden = self.pair_mlp(pair_features) existence_logits = self.existence_head(pair_hidden) positive_logits = self.positive_relation_head(pair_hidden) packed_logits = torch.cat([existence_logits, positive_logits], dim=-1) loss = None if labels is not None: existence_targets = (labels != self.no_relation_id).long() existence_loss = nn.functional.cross_entropy(existence_logits, existence_targets) positive_mask = existence_targets == 1 if positive_mask.any(): positive_targets = self.original_id_to_positive_id[labels[positive_mask]] positive_loss = nn.functional.cross_entropy( positive_logits[positive_mask], positive_targets, ) else: positive_loss = existence_loss.new_zeros(()) loss = ( self.existence_loss_weight * existence_loss + self.positive_relation_loss_weight * positive_loss ) result = {"logits": packed_logits} if loss is not None: result["loss"] = loss return result